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

import duckdb
import polars as pl
import pyarrow as pa
import pyarrow.ipc as ipc
import pyarrow.parquet as pq


def table_bytes(table):
    table = table.cast(pa.schema([(name, pa.int64()) for name in table.column_names])).combine_chunks()
    sink = pa.BufferOutputStream()
    with ipc.new_stream(sink, table.schema) as writer:
        writer.write_table(table, max_chunksize=65536)
    return sink.getvalue().to_pybytes()


def digest(data):
    return hashlib.sha256(data).hexdigest()


def operators(profile, path='0'):
    items = []
    if 'operator_name' in profile:
        items.append({'path': path, 'name': profile['operator_name'], 'type': profile['operator_type'],
                      'outputRows': profile['operator_cardinality'], 'scannedRows': profile['operator_rows_scanned'],
                      'extra': profile['extra_info']})
    for index, child in enumerate(profile.get('children', [])):
        items.extend(operators(child, f'{path}.{index}'))
    return items


def signature(run):
    return [{key: item[key] for key in ['path', 'name', 'type', 'outputRows', 'scannedRows', 'extra']} for item in run['operators']]


def verify_observations(runs):
    by_id = {run['id']: run for run in runs}
    assert by_id['stats']['result'] == by_id['missing']['result'] == by_id['cast']['result']
    assert by_id['stats-stale']['result'] == by_id['stats-refreshed']['result']
    assert by_id['join-auto']['result'] == by_id['join-written']['result']
    assert by_id['stats']['bytesRead'] < by_id['missing']['bytesRead']
    assert by_id['stats']['bytesRead'] < by_id['cast']['bytesRead']
    scans = {name: [op for op in run['operators'] if op['type'] == 'TABLE_SCAN'][0] for name, run in by_id.items()}
    assert scans['independent']['outputRows'] == 100 and scans['correlated']['outputRows'] == 0
    assert scans['independent']['extra']['Estimated Cardinality'] == scans['correlated']['extra']['Estimated Cardinality']
    assert scans['uniform']['outputRows'] == 100 and scans['skew']['outputRows'] == 9505
    assert scans['stats-stale']['outputRows'] == scans['stats-refreshed']['outputRows'] == 1
    assert int(scans['stats-stale']['extra']['Estimated Cardinality']) > int(scans['stats-refreshed']['extra']['Estimated Cardinality'])
    joins = {name: [op for op in by_id[name]['operators'] if op['type'] == 'HASH_JOIN'] for name in ['join-auto', 'join-written']}
    assert [op['outputRows'] for op in joins['join-auto']] == [1000, 100]
    assert [op['outputRows'] for op in joins['join-written']] == [1000, 1000]


parser = argparse.ArgumentParser()
parser.add_argument('--verify', action='store_true')
parser.add_argument('--output', type=Path, default=Path('explorer/static/engines/counterexamples'))
args = parser.parse_args()
output = args.output.resolve()
source_sql = {
    'correlated': 'SELECT i%100 AS a, i%100 AS b FROM range(10000) t(i)',
    'independent': 'SELECT i%100 AS a, (i//100)%100 AS b FROM range(10000) t(i)',
    'uniform': 'SELECT i AS id, i%100 AS key FROM range(10000) t(i)',
    'skew': 'SELECT i AS id, CASE WHEN i<9500 THEN 0 ELSE i%100 END AS key FROM range(10000) t(i)',
    'fact': 'SELECT i AS id, i%10000 AS a, i%100 AS b FROM range(100000) t(i)',
    'wide': 'SELECT i%100 AS b FROM range(1000) t(i)',
    'selective': 'SELECT i AS a, CASE WHEN i<10 THEN 1 ELSE 0 END AS keep FROM range(10000) t(i)',
}
ids = pl.DataFrame({'id': pl.int_range(0, 10000, eager=True)})
frames = {
    'correlated': ids.select((pl.col('id') % 100).alias('a'), (pl.col('id') % 100).alias('b')),
    'independent': ids.select((pl.col('id') % 100).alias('a'), (pl.col('id') // 100 % 100).alias('b')),
    'uniform': ids.with_columns((pl.col('id') % 100).alias('key')),
    'skew': ids.with_columns(pl.when(pl.col('id') < 9500).then(0).otherwise(pl.col('id') % 100).alias('key')),
    'fact': pl.DataFrame({'id': pl.int_range(0, 100000, eager=True)}).with_columns((pl.col('id') % 10000).alias('a'), (pl.col('id') % 100).alias('b')),
    'wide': pl.DataFrame({'b': pl.int_range(0, 1000, eager=True) % 100}),
    'selective': pl.DataFrame({'a': pl.int_range(0, 10000, eager=True)}).with_columns((pl.col('a') < 10).cast(pl.Int64).alias('keep')),
}
for column in ['a', 'b']:
    assert frames['correlated'].group_by(column).len().sort(column).equals(frames['independent'].group_by(column).len().sort(column))
parquet_frame = frames['fact'].select('id').with_columns((pl.col('id') * 17).alias('value'))
join_sql = 'SELECT count(*) AS n, sum(f.id)::BIGINT AS total FROM fact f JOIN wide w USING(b) JOIN selective s USING(a) WHERE s.keep=1'
queries = {
    'stats': "SELECT sum(value)::BIGINT AS total FROM read_parquet('stats.parquet') WHERE id=43",
    'missing': "SELECT sum(value)::BIGINT AS total FROM read_parquet('missing.parquet') WHERE id=43",
    'cast': "SELECT sum(value)::BIGINT AS total FROM read_parquet('stats.parquet') WHERE CAST(id AS VARCHAR)='43'",
    'independent': 'SELECT count(*) AS n FROM independent WHERE a<10 AND b>=90',
    'correlated': 'SELECT count(*) AS n FROM correlated WHERE a<10 AND b>=90',
    'uniform': 'SELECT count(*) AS n FROM uniform WHERE key=0',
    'skew': 'SELECT count(*) AS n FROM skew WHERE key=0',
    'stats-before': 'SELECT count(*) AS n, sum(id)::BIGINT AS total FROM uniform WHERE key=43',
    'stats-stale': 'SELECT count(*) AS n, sum(id)::BIGINT AS total FROM uniform WHERE key=43',
    'stats-refreshed': 'SELECT count(*) AS n, sum(id)::BIGINT AS total FROM uniform WHERE key=43',
    'join-auto': join_sql, 'join-written': join_sql,
}
expected = {name: parquet_frame.filter(pl.col('id') == 43).select(pl.col('value').sum().alias('total')) for name in ['stats', 'missing', 'cast']}
for name in ['independent', 'correlated']:
    expected[name] = frames[name].filter((pl.col('a') < 10) & (pl.col('b') >= 90)).select(pl.len().alias('n'))
for name in ['uniform', 'skew']:
    expected[name] = frames[name].filter(pl.col('key') == 0).select(pl.len().alias('n'))
expected['stats-before'] = frames['uniform'].filter(pl.col('key') == 43).select(pl.len().alias('n'), pl.col('id').sum().alias('total'))
for name in ['stats-stale', 'stats-refreshed']:
    expected[name] = frames['uniform'].with_columns(pl.col('id').alias('key')).filter(pl.col('key') == 43).select(pl.len().alias('n'), pl.col('id').sum().alias('total'))
joined = frames['fact'].join(frames['wide'], on='b').join(frames['selective'].filter(pl.col('keep') == 1), on='a')
for name in ['join-auto', 'join-written']:
    expected[name] = joined.select(pl.len().alias('n'), pl.col('id').sum().alias('total'))

artifacts = {}
runs = []
with tempfile.TemporaryDirectory(prefix='columnar-counterexamples-') as work:
    root = Path(work)
    parquet_metadata = {}
    for name, statistics in [('stats', True), ('missing', False)]:
        path = root / (name + '.parquet')
        pq.write_table(parquet_frame.to_arrow(), path, row_group_size=4096, write_statistics=statistics, compression='NONE')
        assert pl.read_parquet(path).equals(parquet_frame)
        metadata = pq.read_metadata(path)
        assert all((metadata.row_group(i).column(j).statistics is not None) == statistics for i in range(metadata.num_row_groups) for j in range(metadata.num_columns))
        assert all(metadata.row_group(i).column(j).compression == 'UNCOMPRESSED' for i in range(metadata.num_row_groups) for j in range(metadata.num_columns))
        parquet_metadata[name] = {'rows': metadata.num_rows, 'groups': metadata.num_row_groups,
                                  'idBounds': [{'min': metadata.row_group(i).column(0).statistics.min, 'max': metadata.row_group(i).column(0).statistics.max} if statistics else None for i in range(metadata.num_row_groups)]}
        artifacts[path.name] = path.read_bytes()
    assert sum(bounds['min'] <= 43 <= bounds['max'] for bounds in parquet_metadata['stats']['idBounds']) == 1
    fixtures = {name: {'sql': source_sql[name], 'rows': frame.height, 'sha256': digest(table_bytes(frame.to_arrow()))} for name, frame in frames.items()}
    for name, sql in queries.items():
        connection = duckdb.connect(config={'threads': 1})
        connection.execute('SET file_search_path=?', [str(root)])
        for table, source in source_sql.items():
            connection.execute(f'CREATE TABLE {table} AS {source}')
            actual = connection.execute(f'SELECT * FROM {table}').fetch_arrow_table()
            assert table_bytes(actual) == table_bytes(frames[table].to_arrow()), table
        setup = []
        if name.startswith('stats-'):
            setup.append('ANALYZE uniform')
            if name != 'stats-before':
                setup.append('UPDATE uniform SET key=id')
            if name == 'stats-refreshed':
                setup.append('ANALYZE uniform')
        if name == 'join-written':
            setup.append("SET disabled_optimizers='join_order,build_side_probe_side'")
        for statement in setup:
            connection.execute(statement)
        connection.execute("SET enable_profiling='json'")
        profile_path = root / (name + '.json')
        connection.execute('SET profiling_output=?', [str(profile_path)])
        result = connection.execute(sql).fetch_arrow_table()
        assert table_bytes(result) == table_bytes(expected[name].to_arrow()), name
        profile = json.loads(profile_path.read_text().replace(str(root), '<workspace>'))
        connection.close()
        artifacts[profile_path.name] = (json.dumps(profile, indent=2) + '\n').encode()
        runs.append({'id': name, 'sql': sql, 'setup': setup, 'result': result.to_pylist(),
                     'resultSha256': digest(table_bytes(result)), 'profile': profile_path.name,
                     'bytesRead': profile['total_bytes_read'], 'operators': operators(profile)})
        print(name, runs[-1]['result'], runs[-1]['bytesRead'], flush=True)
verify_observations(runs)
proof = {'schemaVersion': 1, 'engine': 'DuckDB ' + duckdb.__version__, 'oracle': 'Polars ' + pl.__version__ + ' / PyArrow ' + pa.__version__,
         'recordedAt': datetime.now(timezone.utc).isoformat(), 'platform': {'system': platform.system(), 'architecture': platform.machine()},
         'configuration': {'threads': 1, 'freshConnectionPerRun': True},
         'normalization': 'Only the temporary workspace path is replaced. Native timings and counters remain in each profile. Input/result digests cover canonical int64 Arrow IPC streams. Byte counters describe this engine execution, not physical disk traffic.',
         'fixtures': fixtures, 'parquetMetadata': parquet_metadata, 'runs': runs,
         'artifacts': {name: {'size': len(data), 'sha256': digest(data)} for name, data in artifacts.items()}}
if args.verify:
    recorded = json.loads((output / 'proof.json').read_text())
    for key in ['schemaVersion', 'engine', 'oracle', 'configuration', 'fixtures', 'parquetMetadata']:
        assert recorded[key] == proof[key], key
    assert [run['id'] for run in recorded['runs']] == list(queries)
    assert set(recorded['artifacts']) == set(artifacts)
    for old, fresh in zip(recorded['runs'], runs, strict=True):
        for key in ['id', 'sql', 'setup', 'result', 'resultSha256', 'profile']:
            assert old[key] == fresh[key], (old['id'], key)
        assert signature(old) == signature(fresh), old['id']
        profile = json.loads((output / old['profile']).read_text())
        assert old['operators'] == operators(profile)
        assert old['bytesRead'] == profile['total_bytes_read']
        assert profile['query_name'] == old['sql'] and profile['rows_returned'] == len(old['result'])
    for name, item in recorded['artifacts'].items():
        data = (output / name).read_bytes()
        assert item == {'size': len(data), 'sha256': digest(data)}, name
        if name.endswith('.parquet'):
            assert data == artifacts[name], name
    verify_observations(recorded['runs'])
    assert (output / 'reproduce.py').read_bytes() == Path(__file__).read_bytes()
    print('Verified all native results, operator estimates and cardinalities, original artifacts and paired observations.')
else:
    assert not (output / 'proof.json').exists(), 'Refusing to replace an existing capture; use a new output directory.'
    output.mkdir(parents=True, exist_ok=True)
    for name, data in artifacts.items():
        (output / name).write_bytes(data)
    (output / 'proof.json').write_text(json.dumps(proof, indent=2) + '\n')
    (output / 'reproduce.py').write_bytes(Path(__file__).read_bytes())
