"""Current-revision clang A/B measurement; no production files are modified."""
import hashlib
import importlib.metadata
import json
import pickle
import platform
import statistics
import subprocess
import sys
import time
from dataclasses import asdict
from pathlib import Path


def dump(path, value):
    path.write_text(json.dumps(value, indent=2, sort_keys=True) + '\n')


def digest(value):
    return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(',', ':')).encode()).hexdigest()


def source_digest(root):
    h = hashlib.sha256()
    files = sorted(p for p in root.rglob('*') if p.is_file())
    for path in files:
        h.update(path.relative_to(root).as_posix().encode() + b'\0')
        h.update(hashlib.sha256(path.read_bytes()).digest())
    return {'files': len(files), 'sha256': h.hexdigest()}


def run(root, mode, index, out):
    from orchestrator.pkg import c_extractor, clang_link
    from orchestrator.pkg.extractor import RepoCodeExtractor
    from orchestrator.pkg.facts import EdgeKind, NodeKind
    from orchestrator.pkg.verify import verify_batch

    assert clang_link.clang_available(), 'Installed wheel required for both modes'
    if mode == 'off':
        clang_link.clang_available = lambda: False
    routed = []
    original_headers = c_extractor.cpp_header_paths
    original_link = clang_link.link_clang
    pending_snapshot = []

    def headers(*args, **kwargs):
        result = original_headers(*args, **kwargs)
        routed.append(result)
        return result

    def link(batch, root, *, pending, report):
        pending_snapshot.extend(pending)
        return original_link(batch, root, pending=pending, report=report)

    c_extractor.cpp_header_paths = headers
    clang_link.link_clang = link
    start = time.perf_counter()
    extractor = RepoCodeExtractor()
    batch = extractor.extract(root)
    elapsed = time.perf_counter() - start
    # All hashing, serialization and verification are outside the extraction timer.
    nodes, edges = batch.nodes, batch.edges
    ids = {n.id for n in nodes}
    grounded = {n.id for n in nodes if n.grounded}
    verification = verify_batch(batch, root)
    node_records = [asdict(n) for n in sorted(nodes, key=lambda n: n.id)]
    edge_records = [asdict(e) for e in sorted(edges, key=lambda e: e.key())]
    pending_records = [asdict(p) for p in sorted(pending_snapshot, key=lambda p: (p.file, p.offset, p.end_offset, p.caller))]
    assert len(routed) == 1
    report = asdict(extractor.clang_report)
    assert report['available'] == (mode == 'on')
    if mode == 'off':
        assert report['resolved'] == report['parsed_tus'] == 0
    record = {
        'mode': mode, 'index': index, 'seconds': elapsed,
        'nodes': len(nodes), 'edges': len(edges),
        'node_sha256': digest(node_records), 'edge_sha256': digest(edge_records),
        'pending_sha256': digest(pending_records),
        'routed_headers': len(routed[0]), 'routed_headers_sha256': digest(sorted(routed[0])),
        'header_types': sum(n.kind is NodeKind.TYPE and n.provenance is not None and n.provenance.file.endswith('.h') and n.language in {'c','cpp'} for n in nodes),
        'calls': sum(e.kind is EdgeKind.CALLS for e in edges),
        'dangling_edges': sum(e.src not in ids or e.dst not in ids for e in edges),
        'clang': report, 'verify': asdict(verification),
        'verification_errors': len(verification.errors), 'verification_warnings': len(verification.warnings),
    }
    dump(out / f'{index}-{mode}.json', record)
    snapshot = out / f'{mode}.pickle'
    if not snapshot.exists():
        snapshot.write_bytes(pickle.dumps(batch))
    print(json.dumps({'repo': root.name, 'mode': mode, 'index': index, 'seconds': round(elapsed,3), 'nodes': len(nodes), 'edges': len(edges), 'resolved': report['resolved'], 'errors':len(verification.errors), 'warnings':len(verification.warnings)}), flush=True)


def compare(out):
    from orchestrator.pkg.facts import EdgeKind, NodeKind
    off = pickle.loads((out / 'off.pickle').read_bytes())
    on = pickle.loads((out / 'on.pickle').read_bytes())
    records = [json.loads(p.read_text()) for p in sorted(out.glob('[1-6]-*.json'))]
    assert len(records) == 6
    added = set(on.edges) - set(off.edges)
    removed = set(off.edges) - set(on.edges)
    off_nodes = {n.id: n for n in off.nodes}
    on_nodes = {n.id: n for n in on.nodes}
    invalid = [e for e in added if e.kind is not EdgeKind.CALLS or any(i not in off_nodes or not off_nodes[i].grounded or off_nodes[i].kind is not NodeKind.FUNCTION for i in (e.src,e.dst))]
    checks = {
        'identical_nodes': off_nodes == on_nodes,
        'identical_node_order': off.nodes == on.nodes,
        'existing_edges_preserved': not removed,
        'added_edges_only_calls_between_preexisting_grounded_functions': not invalid,
        'identical_routing_all_runs': len({r['routed_headers_sha256'] for r in records}) == 1,
        'identical_pending_all_runs': len({r['pending_sha256'] for r in records}) == 1,
        'identical_nodes_all_runs': len({r['node_sha256'] for r in records}) == 1,
        'identical_verification_all_runs': all(r['verify'] == records[0]['verify'] for r in records),
    }
    timing = {}
    for mode in ('off', 'on'):
        group = [r for r in records if r['mode'] == mode]
        checks[f'{mode}_edges_repeatable'] = len({r['edge_sha256'] for r in group}) == 1
        checks[f'{mode}_clang_report_repeatable'] = all(r['clang'] == group[0]['clang'] for r in group)
        times = [r['seconds'] for r in group]
        timing[mode] = {'runs': times, 'median': statistics.median(times), 'min':min(times), 'max':max(times)}
    timing['median_delta_seconds'] = timing['on']['median'] - timing['off']['median']
    timing['median_ratio'] = timing['on']['median'] / timing['off']['median']
    added_records = []
    for e in sorted(added, key=lambda e: e.key()):
        added_records.append({'edge': asdict(e), 'caller_declaration': asdict(off_nodes[e.src]), 'target_declaration': asdict(off_nodes[e.dst])})
    dump(out / 'added-edges.json', added_records)
    dump(out / 'comparison.json', {'checks':checks,'timing':timing,'added_edges':len(added),'removed_edges':len(removed),'invalid_added_edges':len(invalid),'runs':records})
    print(json.dumps({'repo': out.name, 'checks':checks,'added':len(added),'removed':len(removed),'timing':timing}), flush=True)
    assert all(checks.values()), checks


if __name__ == '__main__':
    if sys.argv[1] == 'run':
        run(Path(sys.argv[2]), sys.argv[3], int(sys.argv[4]), Path(sys.argv[5]))
    elif sys.argv[1] == 'compare':
        compare(Path(sys.argv[2]))
    else:
        output = Path(sys.argv[2]).resolve()
        output.mkdir(exist_ok=True, parents=True)
        repos = [
            ('tinyxml2', Path('/tmp/spine-clang-ab-tinyxml2'), '8224e427b655b83dae5e2298f1e6919523a78737'),
            ('opencv', Path('/tmp/spine-clang-ab-opencv'), 'b4c5ec4042f097e2a5b386b9d413ec7333d0a184'),
        ]
        metadata = {
            'spine_commit':subprocess.check_output(['git','rev-parse','HEAD'], text=True).strip(),
            'python':sys.version, 'platform':platform.platform(), 'machine':platform.machine(),
            'packages':{d.metadata['Name']:d.version for d in importlib.metadata.distributions()},
            'order':['off','on','on','off','off','on'],
            'method':'Fresh subprocess per extraction; only clang_available is forced false in off mode. Hashing and verify_batch outside timer. Source hashing pre-reads every input before the first run. No extraction cache. No parallel benchmark or test suite.',
            'repositories':{},
        }
        for name, root, commit in repos:
            metadata['repositories'][name] = {'commit':commit,'before':source_digest(root)}
        dump(output / 'environment.json', metadata)
        for name, root, commit in repos:
            out = output / name
            out.mkdir(exist_ok=False)
            for index, mode in enumerate(metadata['order'], 1):
                subprocess.run([sys.executable,__file__,'run',str(root),mode,str(index),str(out)],check=True)
            subprocess.run([sys.executable,__file__,'compare',str(out)],check=True)
            metadata['repositories'][name]['after'] = source_digest(root)
            assert metadata['repositories'][name]['before'] == metadata['repositories'][name]['after']
            dump(output / 'environment.json', metadata)
        print('COMPLETE: all current-revision A/B checks passed', flush=True)
