# /// script
# requires-python = ">=3.12,<3.13"
# dependencies = ["pyarrow==25.0.1", "duckdb==1.4.4"]
# ///
import argparse
import hashlib
import json
from pathlib import Path

import duckdb
import pyarrow as pa
import pyarrow.compute as pc
import pyarrow.parquet as pq

parser = argparse.ArgumentParser()
parser.add_argument('--source-root', type=Path, default=Path('explorer/static'))
parser.add_argument('--output', type=Path)
parser.add_argument('--module-output', type=Path)
parser.add_argument('--verify', action='store_true')
args = parser.parse_args()
static = args.source_root.resolve()
target = args.output or static / 'table-snapshots'
verify = args.verify
artifacts = {}
entries = []
con = duckdb.connect()


def proof(path):
    payload = (static / path).read_bytes()
    return json.loads(payload), hashlib.sha256(payload).hexdigest()


def export(key, family, label, version, metadata, source, source_hash, rows, href, renamed=False):
    table = pa.Table.from_pylist(rows).select(['observation_id', 'station', 'observed_at', 'temperature_c'])
    table = table.set_column(0, 'observation_id', table['observation_id'].cast(pa.int64()))
    timestamps = pc.cast(pc.replace_substring(table['observed_at'], 'Z', ''), pa.timestamp('us'))
    if family.startswith('delta'):
        timestamps = timestamps.cast(pa.timestamp('us', tz='UTC'))
    table = table.set_column(2, 'observed_at', timestamps)
    table = table.set_column(3, 'temperature_c', table['temperature_c'].cast(pa.float64()))
    if renamed:
        table = table.rename_columns(['observation_id', 'station', 'observed_at', 'air_temperature_c'])
    table = table.sort_by([('observation_id', 'ascending'), (table.column_names[3], 'ascending')])
    table = table.replace_schema_metadata({
        'columnar.kind': 'derived-native-snapshot-result',
        'columnar.source': source,
        'columnar.source_sha256': source_hash,
        'columnar.version': version,
    })
    sink = pa.BufferOutputStream()
    pq.write_table(table, sink, compression='snappy', use_dictionary=False, write_page_checksum=True)
    payload = sink.getvalue().to_pybytes()
    assert pq.read_table(pa.BufferReader(payload)).equals(table)
    con.register('expected_snapshot', table)
    projection = f'observation_id, station, epoch_us(observed_at), "{table.column_names[3]}"'
    expected = con.execute(f'SELECT {projection} FROM expected_snapshot ORDER BY observation_id, 4').fetchall()
    name = key + '.parquet'
    artifacts[name] = payload
    entries.append({
        'id': key, 'family': family, 'label': label, 'version': version, 'metadata': metadata,
        'source': '/' + source, 'sourceSha256': source_hash, 'href': href,
        'url': '/table-snapshots/' + name, 'size': len(payload),
        'sha256': hashlib.sha256(payload).hexdigest(), 'rows': table.num_rows,
        'columns': [{'name': field.name, 'type': str(field.type)} for field in table.schema],
        'ids': table['observation_id'].to_pylist(),
        'sql': 'SELECT * FROM data ORDER BY observation_id, 4',
    })
    if verify:
        assert (target / name).read_bytes() == payload, name
        actual = con.execute(f'SELECT {projection} FROM read_parquet(?) ORDER BY observation_id, 4', [str(target / name)]).fetchall()
        assert actual == expected, name
        temperature = table.column_names[3]
        narrow = con.execute(f'SELECT "{temperature}" FROM read_parquet(?) ORDER BY "{temperature}"', [str(target / name)]).fetchall()
        assert narrow == con.execute(f'SELECT "{temperature}" FROM expected_snapshot ORDER BY "{temperature}"').fetchall(), name


iceberg, iceberg_hash = proof('iceberg/weather/index.json')
for index, snapshot in enumerate(iceberg['snapshots']):
    for renamed in ([False, True] if index == len(iceberg['snapshots']) - 1 else [False]):
        metadata = iceberg['metadataVersions'][3]['path'] if renamed else snapshot['metadataPath']
        export(f'iceberg-weather-{index}' + ('-renamed' if renamed else ''), 'iceberg-weather',
               snapshot['label'] + (' · renamed schema' if renamed else ''), snapshot['id'], metadata,
               'iceberg/weather/index.json', iceberg_hash, snapshot['expectedRows'],
               f"/iceberg/explained?snapshot={snapshot['id']}" + ('&schema=renamed' if renamed else '') + '#membership', renamed)

delta, delta_hash = proof('delta/weather/index.json')
for version in delta['versions']:
    export(f"delta-weather-{version['version']}", 'delta-weather', f"Version {version['version']}",
           str(version['version']), f"_delta_log/{version['version']:020}.json", 'delta/weather/index.json',
           delta_hash, version['rows'], f"/delta/explained?version={version['version']}#snapshot")

deletes, deletes_hash = proof('iceberg/deletes/proof.json')
for case in deletes['cases']:
    for index, snapshot in enumerate(case['snapshots']):
        export(f"iceberg-{case['key']}-{index}", 'iceberg-deletes', case['key'] + ' · ' + snapshot['label'],
               snapshot['id'], snapshot['metadata'], 'iceberg/deletes/proof.json', deletes_hash,
               snapshot['rows'], f"/iceberg/explained?deleteCase={case['key']}&deleteStep={index}#delete-files")

vectors, vectors_hash = proof('delta/native/proof.json')
source_rows = pa.Table.from_pylist(delta['versions'][1]['rows'])
for version in vectors['versions']:
    selected = source_rows.filter(pc.is_in(source_rows['observation_id'], value_set=pa.array(version['ids'], type=pa.int64())))
    assert sorted(selected['observation_id'].to_pylist()) == version['ids']
    export(f"delta-vectors-{version['version']}", 'delta-vectors', f"Deletion-vector version {version['version']}",
           str(version['version']), f"table/_delta_log/{version['version']:020}.json", 'delta/native/proof.json',
           vectors_hash, selected.to_pylist(), f"/delta/explained?dvVersion={version['version']}#delta-native-vectors")

manifest = {
    'kind': 'derived-native-snapshot-results', 'writer': 'PyArrow 25.0.1', 'reader': 'DuckDB 1.4.4',
    'scope': 'Parquet exports of recorded native logical results. Native verifiers validate the source proofs; the browser queries these exports, not a live table.',
    'entries': entries,
}
artifacts['manifest.json'] = (json.dumps(manifest, indent=2) + '\n').encode()
artifacts['reproduce.py'] = Path(__file__).read_bytes()
if args.module_output:
    if verify:
        assert args.module_output.read_bytes() == artifacts['manifest.json']
    else:
        args.module_output.parent.mkdir(parents=True, exist_ok=True)
        args.module_output.write_bytes(artifacts['manifest.json'])
if verify:
    for name, payload in artifacts.items():
        assert (target / name).read_bytes() == payload, name
else:
    target.mkdir(parents=True, exist_ok=True)
    for name, payload in artifacts.items():
        (target / name).write_bytes(payload)
print(f'{"Verified" if verify else "Wrote"} {len(entries)} derived snapshot exports, source hashes, exact records and narrow projections.')
