# /// script
# requires-python = ">=3.12,<3.13"
# dependencies = ["pyspark==3.5.6", "delta-spark==3.3.2", "deltalake==1.6.3", "pyarrow==25.0.1", "pyroaring==1.0.3", "pyzmq==27.1.0"]
# ///
import argparse
import hashlib
import json
import os
import shutil
import struct
import sys
import tempfile
import uuid
import zlib
from pathlib import Path

import pyarrow.parquet as pq
from delta import configure_spark_with_delta_pip
from deltalake import DeltaTable, write_deltalake
from deltalake.exceptions import CommitFailedError
from pyroaring import BitMap
from pyspark.sql import SparkSession
from zmq.utils.z85 import decode

ROOT = Path.cwd()
parser = argparse.ArgumentParser()
parser.add_argument('--verify', action='store_true')
parser.add_argument('--source', type=Path, default=ROOT/'explorer/static/delta/weather')
parser.add_argument('--output', type=Path, default=ROOT/'explorer/static/delta/native')
args = parser.parse_args()
source = args.source.resolve()
output = args.output.resolve()
workspace = Path(tempfile.mkdtemp(prefix='columnar-delta-native-'))


def actions(path, version):
    return [json.loads(line) for line in (path/'_delta_log'/f'{version:020}.json').read_text().splitlines()]


def clone(name):
    path = workspace/name
    (path/'_delta_log').mkdir(parents=True)
    for version in range(2):
        shutil.copy2(source/'_delta_log'/f'{version:020}.json', path/'_delta_log'/f'{version:020}.json')
        for action in actions(source, version):
            if 'add' in action:
                name = action['add']['path']
                shutil.copy2(source/name, path/name)
    return path


def ids(table):
    return sorted(table.column('observation_id').to_pylist())


def bitmap(path, descriptor):
    assert descriptor['storageType'] == 'u'
    value = descriptor['pathOrInlineDv']
    name = Path(value[:-20])/f'deletion_vector_{uuid.UUID(bytes=decode(value[-20:]))}.bin'
    raw = (path/name).read_bytes()
    offset, size = descriptor['offset'], descriptor['sizeInBytes']
    assert raw[0] == 1 and struct.unpack_from('>I', raw, offset)[0] == size
    data = raw[offset+4:offset+4+size]
    assert struct.unpack_from('>I', raw, offset+4+size)[0] == zlib.crc32(data)
    assert struct.unpack_from('<I', data)[0] == 1681511377
    buckets = struct.unpack_from('<Q', data, 4)[0]
    assert buckets == 1
    high = struct.unpack_from('<I', data, 12)[0]
    positions = [(high << 32) + low for low in BitMap.deserialize(data[16:])]
    assert len(positions) == descriptor['cardinality']
    return {'path': str(name), 'hex': raw.hex(), 'positions': positions, 'offset': offset, 'size': size, 'fileBytes': len(raw)}


def races():
    results = []
    for kind in ['overlap', 'append', 'append-pair']:
        path = clone(kind)
        a, b = DeltaTable(path), DeltaTable(path)
        assert a.version() == b.version() == 1
        pinned = ids(b.to_pyarrow_table())
        second = pq.read_table(source/next(action['add']['path'] for action in actions(source,1) if 'add' in action))
        if kind == 'append-pair':
            write_deltalake(a, second.slice(0,1), mode='append')
        else:
            a.delete('observation_id = 13')
        before = set(path.glob('*.parquet'))
        error = None
        if kind == 'overlap':
            try:
                b.delete('observation_id = 14')
            except CommitFailedError as cause:
                error = str(cause)
            assert error is not None, 'Expected stale replacement to conflict'
            assert DeltaTable(path).version() == 2
            candidate = sorted(set().union(*(set(ids(pq.read_table(file))) for file in set(path.glob('*.parquet'))-before)))
            assert 13 in candidate and 14 not in candidate
            DeltaTable(path).delete('observation_id = 14')
            expected = [i for i in range(1,19) if i not in [13,14]]
        else:
            row = second.slice(1 if kind == 'append-pair' else 0,1)
            if kind == 'append':
                try:
                    write_deltalake(b, row, mode='append')
                except CommitFailedError as cause:
                    error = str(cause)
                assert error is not None, 'Pinned client rejects stale append after removal'
                assert DeltaTable(path).version() == 2
                write_deltalake(DeltaTable(path), row, mode='append')
            else:
                write_deltalake(b, row, mode='append')
            candidate = [14] if kind == 'append-pair' else [13]
            expected = sorted(list(range(1,19)) + ([13,14] if kind == 'append-pair' else []))
        actual = ids(DeltaTable(path).to_pyarrow_table())
        assert actual == expected
        assert pinned == list(range(1,19))
        commits = [actions(path,i) for i in [2,3]]
        results.append({'kind': kind, 'readVersion':1, 'error':error, 'privateCandidateIds':candidate, 'afterA':sorted(list(range(1,19))+[13]) if kind=='append-pair' else [i for i in range(1,19) if i != 13], 'resultIds':actual, 'commits':commits})
    return results


os.environ['PYSPARK_PYTHON'] = sys.executable
builder = (SparkSession.builder.master('local[2]').appName('Columnar native Delta depth')
    .config('spark.sql.extensions','io.delta.sql.DeltaSparkSessionExtension')
    .config('spark.sql.catalog.spark_catalog','org.apache.spark.sql.delta.catalog.DeltaCatalog')
    .config('spark.sql.session.timeZone','UTC').config('spark.ui.enabled','false')
    .config('spark.sql.shuffle.partitions','2').config('spark.databricks.delta.snapshotPartitions','2'))
spark = configure_spark_with_delta_pip(builder).getOrCreate()
spark.sparkContext.setLogLevel('ERROR')
path = output/'table' if args.verify else clone('dv')
if not args.verify:
    assert not output.exists(), 'Inspect the existing bundle before choosing a new output directory'
    spark.sql(f"ALTER TABLE delta.`{path}` SET TBLPROPERTIES ('delta.enableDeletionVectors' = 'true')")
    for observation in [13,16]:
        spark.sql(f'DELETE FROM delta.`{path}` WHERE observation_id = {observation}')
base_file = next(a['add']['path'] for a in actions(source,1) if 'add' in a)
physical = pq.read_table(path/base_file).to_pylist()
assert [row['observation_id'] for row in physical] == list(range(13,19))
assert (path/base_file).read_bytes() == (source/base_file).read_bytes()
versions = []
for version in [2,3,4]:
    expected = [i for i in range(1,19) if i not in ([] if version==2 else [13] if version==3 else [13,16])]
    rows = spark.read.format('delta').option('versionAsOf',version).load(str(path)).orderBy('observation_id').collect()
    assert [r.observation_id for r in rows] == expected
    original = {r.observation_id:r for r in spark.read.format('delta').option('versionAsOf',1).load(str(source)).collect()}
    assert all(r == original[r.observation_id] for r in rows)
    commit = actions(path,version)
    add = next((a['add'] for a in commit if 'add' in a),None)
    dv = bitmap(path,add['deletionVector']) if add else None
    if dv:
        assert dv['positions'] == ([0] if version==3 else [0,3])
        assert [r['observation_id'] for i,r in enumerate(physical) if i not in dv['positions']] == [i for i in expected if i>=13]
    warm = spark.read.format('delta').option('versionAsOf',version).load(str(path)).where('temperature_c >= 20').orderBy('observation_id').select('observation_id').collect()
    assert [row.observation_id for row in warm] == ([16,17,18] if version<4 else [17,18])
    versions.append({'version':version,'ids':expected,'warmIds':[row.observation_id for row in warm],'actions':commit,'bitmap':dv})
assert versions[1]['ids'] == [r.observation_id for r in spark.read.format('delta').option('versionAsOf',2).load(str(source)).orderBy('observation_id').collect()]
concurrency = races()
for race in concurrency:
    table = output/race['kind'] if args.verify else workspace/race['kind']
    for version, expected in [(1,list(range(1,19))),(2,race['afterA']),(3,race['resultIds'])]:
        rows = spark.read.format('delta').option('versionAsOf',version).load(str(table)).orderBy('observation_id').collect()
        assert [row.observation_id for row in rows] == expected
        assert all(row == original[row.observation_id] for row in rows)
spark.stop()
proof = {'writer':'Delta Spark 3.3.2 / Spark 3.5.6', 'concurrencyWriter':'delta-rs 1.6.3', 'baseFile':base_file, 'physicalIds':list(range(13,19)), 'physicalTemperatures':[row['temperature_c'] for row in physical], 'versions':versions, 'races':concurrency}
if args.verify:
    recorded = json.loads((output/'proof.json').read_text())
    assert versions == recorded['versions']
    assert [row['temperature_c'] for row in physical] == recorded['physicalTemperatures']
    for actual, expected in zip(concurrency,recorded['races'],strict=True):
        for key in ['kind','readVersion','error','privateCandidateIds','afterA','resultIds']:
            assert actual[key] == expected[key], key
    for artifact in recorded['artifacts']:
        raw=(output/artifact['path']).read_bytes()
        assert len(raw)==artifact['size'] and hashlib.sha256(raw).hexdigest()==artifact['sha256']
    print('Verified original DV bytes, CRC, independent Roaring decode, all Spark values, and all three native writer schedules and unseen-predicate answers.')
else:
    output.mkdir(parents=True)
    shutil.copytree(path,output/'table',ignore=shutil.ignore_patterns('.*'))
    for kind in ['overlap','append','append-pair']:
        shutil.copytree(workspace/kind,output/kind,ignore=shutil.ignore_patterns('.*'))
    artifacts=[]
    for file in sorted(output.rglob('*')):
        if file.is_file():
            raw=file.read_bytes();artifacts.append({'path':str(file.relative_to(output)), 'size':len(raw),'sha256':hashlib.sha256(raw).hexdigest()})
    proof['artifacts']=artifacts
    (output/'proof.json').write_text(json.dumps(proof,indent=2)+'\n')
    print('Generated',output)
