# /// script
# requires-python = ">=3.12,<3.13"
# dependencies = ["pyspark==3.5.6", "delta-spark==3.3.2", "pyarrow==25.0.1", "polars==1.44.1", "fastavro==1.12.1"]
# ///
import argparse
import hashlib
import importlib.metadata
import json
import os
import re
import shutil
import sys
import tempfile
import zipfile
from datetime import datetime, timezone
from pathlib import Path
from urllib.parse import unquote, urlparse

import fastavro
import polars as pl
import pyarrow as pa
import pyarrow.parquet as pq
from delta import configure_spark_with_delta_pip
from pyspark.sql import SparkSession


def hash_file(path):
    return hashlib.sha256(path.read_bytes()).hexdigest()


def path_of(uri):
    return Path(unquote(urlparse(str(uri)).path))


def delta_state(directory, version):
    files = {}
    metadata = None
    protocol = None
    commits = []
    for index in range(version + 1):
        path = directory / '_delta_log' / f'{index:020}.json'
        commits.append(path)
        for action in map(json.loads, path.read_text().splitlines()):
            if 'remove' in action:
                del files[action['remove']['path']]
            if 'add' in action:
                files[action['add']['path']] = action['add']
            if 'metaData' in action:
                metadata = action['metaData']
            if 'protocol' in action:
                protocol = action['protocol']
    assert metadata is not None and protocol is not None
    return files, metadata, protocol, commits


def original_rows(directory, step):
    rows = []
    for item in step['files']:
        path = directory / item['path']
        assert path.stat().st_size == item['size'] and hash_file(path) == item['sha256']
        table = pq.read_table(path)
        columns = {}
        for field in step['fields']:
            physical = field.get('physicalName', field['name'])
            if step['family'] == 'iceberg':
                matches = [column.name for column in table.schema if column.metadata and column.metadata.get(b'PARQUET:field_id') == str(field['id']).encode()]
                assert len(matches) == 1, field
                physical = matches[0]
            columns[field['name']] = table[physical]
        rows.extend(pa.table(columns).to_pylist())
    return sorted(rows, key=lambda row: row['id'])


def check_originals(directory, proof):
    for name, item in proof['artifacts'].items():
        path = directory / name
        assert path.stat().st_size == item['size'] and hash_file(path) == item['sha256'], name
    for step in proof['steps']:
        assert original_rows(directory, step) == step['rows'], (step['family'], step['operation'])
        if step['family'] == 'delta':
            files, metadata, protocol, commits = delta_state(directory / 'delta', int(step['version']))
            assert set('delta/' + path for path in files) == {item['path'] for item in step['files']}
            assert protocol == step['protocol']
            fields = json.loads(metadata['schemaString'])['fields']
            assert step['fields'] == [{'name': field['name'], 'id': field['metadata']['delta.columnMapping.id'], 'physicalName': field['metadata']['delta.columnMapping.physicalName']} for field in fields]
            assert [str(path.relative_to(directory)) for path in commits] == step['dependencies']
        else:
            metadata = json.loads((directory / step['root']).read_text())
            snapshot = next(item for item in metadata['snapshots'] if str(item['snapshot-id']) == step['version'])
            source_root = proof['originalLocations']['iceberg']
            def mapped(uri):
                path = path_of(uri)
                return directory / 'iceberg' / path.relative_to(path_of(source_root))
            manifest_list = mapped(snapshot['manifest-list'])
            dependencies = [step['root'], str(manifest_list.relative_to(directory))]
            active = set()
            with manifest_list.open('rb') as stream:
                manifests = list(fastavro.reader(stream))
            for manifest in manifests:
                path = mapped(manifest['manifest_path'])
                dependencies.append(str(path.relative_to(directory)))
                with path.open('rb') as stream:
                    for entry in fastavro.reader(stream):
                        if entry['status'] in [0, 1]:
                            assert entry['data_file']['content'] == 0
                            active.add(str(mapped(entry['data_file']['file_path']).relative_to(directory)))
            assert active == {item['path'] for item in step['files']}
            assert sorted(dependencies) == sorted(step['dependencies'])


def check_transitions(steps):
    for family in ['iceberg', 'delta']:
        selected = {step['operation']: step for step in steps if step['family'] == family}
        assert len(selected['compact']['files']) < len(selected['update-delete']['files'])
        assert selected['compact']['files'] == selected['rename']['files']
        assert [field['id'] for field in selected['compact']['fields']] == [field['id'] for field in selected['rename']['fields']]
        assert selected['retained']['rows'] == selected['append']['rows']
        assert selected['retained']['files'] == selected['append']['files']
        if family == 'delta':
            assert [field['physicalName'] for field in selected['compact']['fields']] == [field['physicalName'] for field in selected['rename']['fields']]
        else:
            assert selected['compact']['version'] == selected['rename']['version']
            assert selected['compact']['root'] != selected['rename']['root']


parser = argparse.ArgumentParser()
parser.add_argument('--verify', action='store_true')
parser.add_argument('--output', type=Path, default=Path('explorer/static/table-history'))
args = parser.parse_args()
output = args.output.resolve()
base = pl.DataFrame({'id': pl.int_range(1, 13, eager=True)}).with_columns(pl.when(pl.col('id') % 2 == 0).then(pl.lit('north')).otherwise(pl.lit('south')).alias('category'), (pl.col('id') * 10).alias('score'))
changed = base.filter(pl.col('id') != 9).with_columns(pl.when(pl.col('id') == 3).then(pl.col('score') + 7).otherwise(pl.col('score')).alias('score'))
expected = {'create': base.head(8), 'append': base, 'update-delete': changed, 'compact': changed, 'rename': changed.rename({'score': 'reading'}), 'retained': base}
steps = []
os.environ['PYSPARK_PYTHON'] = sys.executable
with tempfile.TemporaryDirectory(prefix='columnar-paired-history-') as work:
    root = Path(work)
    evidence = root / 'evidence'
    builder = (SparkSession.builder.master('local[1]').appName('Columnar paired native histories')
        .config('spark.sql.extensions', 'io.delta.sql.DeltaSparkSessionExtension,org.apache.iceberg.spark.extensions.IcebergSparkSessionExtensions')
        .config('spark.sql.catalog.spark_catalog', 'org.apache.spark.sql.delta.catalog.DeltaCatalog')
        .config('spark.sql.catalog.lab', 'org.apache.iceberg.spark.SparkCatalog')
        .config('spark.sql.catalog.lab.type', 'hadoop')
        .config('spark.sql.catalog.lab.warehouse', (root / 'warehouse').as_uri())
        .config('spark.sql.shuffle.partitions', '1').config('spark.databricks.delta.snapshotPartitions', '1')
        .config('spark.sql.session.timeZone', 'UTC').config('spark.ui.enabled', 'false'))
    spark = configure_spark_with_delta_pip(builder, extra_packages=['org.apache.iceberg:iceberg-spark-runtime-3.5_2.12:1.9.2']).getOrCreate()
    spark.sparkContext.setLogLevel('ERROR')
    assert spark.version == '3.5.6'
    assert importlib.metadata.version('delta-spark') == '3.3.2'
    assert str(spark._jvm.org.apache.iceberg.IcebergBuild.version()) == '1.9.2'
    spark.sql('CREATE NAMESPACE lab.default')
    locations = {'iceberg': root / 'warehouse/default/history', 'delta': root / 'delta'}
    refs = {'iceberg': 'lab.default.history', 'delta': f'delta.`{locations["delta"]}`'}
    spark.sql(f"CREATE TABLE {refs['iceberg']} (id BIGINT, category STRING, score BIGINT) USING iceberg TBLPROPERTIES ('format-version'='2', 'write.delete.mode'='copy-on-write', 'write.update.mode'='copy-on-write')")
    spark.sql(f"CREATE TABLE {refs['delta']} (id BIGINT, category STRING, score BIGINT) USING delta TBLPROPERTIES ('delta.columnMapping.mode'='name', 'delta.minReaderVersion'='2', 'delta.minWriterVersion'='5')")
    tables = spark._jvm.org.apache.iceberg.hadoop.HadoopTables(spark.sparkContext._jsc.hadoopConfiguration())
    table = tables.load(locations['iceberg'].as_uri())
    originals = {family: path.as_uri() for family, path in locations.items()}

    def reference(path, family):
        return family + '/' + str(path.relative_to(locations[family]))

    def capture(family, operation, commands, retained=None):
        table.refresh()
        version = str(table.currentSnapshot().snapshotId()) if family == 'iceberg' else str(max(int(path.stem) for path in (locations['delta'] / '_delta_log').glob('*.json')))
        if retained is not None:
            version = retained['version']
        query = f'SELECT * FROM {refs[family]}' + (f' VERSION AS OF {version}' if retained is not None else '') + ' ORDER BY id'
        frame = spark.sql(query)
        rows = [row.asDict() for row in frame.collect()]
        assert rows == expected[operation].to_dicts(), (family, operation, rows)
        active = []
        if family == 'iceberg':
            snapshot = table.snapshot(int(version))
            tasks = table.newScan().useSnapshot(int(version)).planFiles()
            iterator = tasks.iterator()
            while iterator.hasNext():
                task = iterator.next()
                assert len(task.deletes()) == 0
                active.append(path_of(task.file().location()))
            tasks.close()
            metadata_path = path_of(table.operations().current().metadataFileLocation())
            metadata = json.loads(metadata_path.read_text())
            schema_id = snapshot.schemaId() if retained is not None else metadata['current-schema-id']
            schema = next(schema for schema in metadata['schemas'] if schema['schema-id'] == schema_id)
            fields = [{'name': field['name'], 'id': field['id']} for field in schema['fields']]
            dependencies = [metadata_path, path_of(snapshot.manifestListLocation()), *[path_of(manifest.path()) for manifest in snapshot.allManifests(table.io())]]
            published = reference(metadata_path, family)
            protocol = {'formatVersion': metadata['format-version']}
        else:
            files, metadata, protocol, dependencies = delta_state(locations[family], int(version))
            active = [locations[family] / path for path in files]
            assert set(path_of(uri).resolve() for uri in frame.inputFiles()) == set(path.resolve() for path in active)
            fields = [{'name': field['name'], 'id': field['metadata']['delta.columnMapping.id'], 'physicalName': field['metadata']['delta.columnMapping.physicalName']} for field in json.loads(metadata['schemaString'])['fields']]
            published = reference(dependencies[-1], family)
        files = [{'path': reference(path, family), 'size': path.stat().st_size, 'sha256': hash_file(path)} for path in sorted(active)]
        step = {'family': family, 'operation': operation, 'version': version, 'root': published,
                'commands': commands, 'query': query.replace(str(root), '<workspace>'), 'rows': rows,
                'fields': fields, 'protocol': protocol, 'files': files,
                'dependencies': [reference(path, family) for path in dependencies]}
        steps.append(step)
        print(family, operation, version, len(rows), len(files), flush=True)
        return step

    retained = {}
    for operation in ['create', 'append', 'update-delete', 'compact', 'rename', 'retained']:
        for family, ref in refs.items():
            commands = []
            if operation in ['create', 'append']:
                ids = range(1, 9) if operation == 'create' else range(9, 13)
                values = ','.join(f"({i},'{['north', 'south'][i % 2]}',{i * 10})" for i in ids)
                commands = [f'INSERT INTO {ref} VALUES {values}']
            elif operation == 'update-delete':
                commands = [f'UPDATE {ref} SET score=score+7 WHERE id=3', f'DELETE FROM {ref} WHERE id=9']
            elif operation == 'compact':
                commands = ["CALL lab.system.rewrite_data_files(table=>'default.history', options=>map('rewrite-all','true'))"] if family == 'iceberg' else [f'OPTIMIZE {ref}']
            elif operation == 'rename':
                commands = [f'ALTER TABLE {ref} RENAME COLUMN score TO reading']
            for command in commands:
                spark.sql(command).collect()
            step = capture(family, operation, [command.replace(str(root), '<workspace>') for command in commands], retained[family] if operation == 'retained' else None)
            if operation == 'append':
                retained[family] = step
    spark.stop()
    check_transitions(steps)
    evidence.mkdir()
    with zipfile.ZipFile(evidence / 'originals.zip', 'w', compression=zipfile.ZIP_DEFLATED) as archive:
        for family, location in locations.items():
            for path in sorted(location.rglob('*')):
                if path.is_file():
                    name = reference(path, family)
                    archive.write(path, name)
                    if not path.name.startswith('.'):
                        destination = evidence / name
                        destination.parent.mkdir(parents=True, exist_ok=True)
                        shutil.copyfile(path, destination)
    artifacts = {str(path.relative_to(evidence)): {'size': path.stat().st_size, 'sha256': hash_file(path)} for path in sorted(evidence.rglob('*')) if path.is_file()}
    proof = {'schemaVersion': 1, 'recordedAt': datetime.now(timezone.utc).isoformat(),
             'writers': {'spark': '3.5.6', 'iceberg': '1.9.2', 'delta': '3.3.2'}, 'oracle': 'Polars ' + pl.__version__ + ' / PyArrow ' + pa.__version__,
             'scope': 'Local Hadoop Iceberg catalog and path-based Delta, copy-on-write operations, Delta name column mapping enabled from table creation. No REST catalog, CDF, merge-on-read or generic browser table execution.',
             'originalLocations': originals, 'steps': steps, 'artifacts': artifacts}
    check_originals(evidence, proof)
    if args.verify:
        recorded = json.loads((output / 'proof.json').read_text())
        assert recorded['schemaVersion'] == 1 and recorded['writers'] == proof['writers'] and recorded['oracle'] == proof['oracle']
        assert len(recorded['steps']) == len(steps)
        for old, fresh in zip(recorded['steps'], steps, strict=True):
            for key in ['family', 'operation', 'rows', 'fields', 'protocol', 'commands']:
                if key == 'fields' and old['family'] == 'delta':
                    assert [{k: v for k, v in field.items() if k != 'physicalName'} for field in old[key]] == [{k: v for k, v in field.items() if k != 'physicalName'} for field in fresh[key]]
                else:
                    assert old[key] == fresh[key], (old['operation'], key)
            assert len(old['files']) == len(fresh['files'])
            assert re.sub(r' VERSION AS OF \d+', ' VERSION AS OF <version>', old['query']) == re.sub(r' VERSION AS OF \d+', ' VERSION AS OF <version>', fresh['query'])
        check_transitions(recorded['steps'])
        check_originals(output, recorded)
        assert (output / 'reproduce.py').read_bytes() == Path(__file__).read_bytes()
        print('Verified both complete histories and retained reads against independent rows and original metadata/data artifacts.')
    else:
        assert not output.exists(), 'Use a new output directory to preserve existing evidence.'
        shutil.copytree(evidence, output)
        (output / 'proof.json').write_text(json.dumps(proof, indent=2) + '\n')
        shutil.copyfile(__file__, output / 'reproduce.py')
