#!/usr/bin/env python3
"""Finite exact-library controls and selected source-guard arithmetic; no source execution."""
from datetime import datetime, timezone
from fractions import Fraction
import hashlib
import itertools
import json
from pathlib import Path
import platform
import statistics
import sys
import time
import unicodedata

import rapidfuzz
from rapidfuzz.distance import Levenshtein

ROOT = Path(__file__).resolve().parents[2]
OUT = ROOT / 'research/prototypes/edit-distance-comparison-2026-10-09-v1'
REPORT = ROOT / 'research/snapshots/2026-10-08-baseline/edit-distance-baseline-validation-2026-10-09-v1.json'


def clog2(n):
    return (n - 1).bit_length() if n > 1 else 0


def parameters(n):
    h = clog2(clog2(n + 2))
    return dict(N=str(n), height_exponent=h, H=1 << h, ell=1 << clog2(max(1, h)))


def dp(a, b):
    row = list(range(len(b) + 1))
    for i, x in enumerate(a, 1):
        nxt = [i]
        for j, y in enumerate(b, 1):
            nxt.append(min(row[j] + 1, nxt[-1] + 1, row[j - 1] + (x != y)))
        row = nxt
    return row[-1]


def jsave(path, value):
    with path.open('x') as stream:
        json.dump(value, stream, indent=2, ensure_ascii=False)
        stream.write('\n')


def main():
    if OUT.exists() or REPORT.exists():
        raise FileExistsError('Preserved comparison already exists')
    OUT.mkdir()
    checks = []

    def check(label, condition):
        checks.append(dict(label=label, passed=bool(condition)))
        if not condition:
            raise AssertionError(label)

    strings = [''.join(p) for n in range(5) for p in itertools.product('ab', repeat=n)]
    for a, b in itertools.product(strings, repeat=2):
        expected = dp(a, b)
        ops = Levenshtein.editops(a, b)
        cutoffs = [Levenshtein.distance(a, b, score_cutoff=k) for k in (0, 1, 2, 4)]
        check(f'binary-pair:{a!r}:{b!r}',
              Levenshtein.distance(a, b) == expected
              and cutoffs == [expected if expected <= k else k + 1 for k in (0, 1, 2, 4)]
              and len(ops) == expected and ops.apply(a, b) == b)

    # A separate multiplication-loop oracle checks the bounded integer logarithm.
    for n in range(1025):
        exponent, bound = 0, 1
        while bound < n:
            exponent, bound = exponent + 1, bound * 2
        check(f'ceil-log2:{n}', clog2(n) == exponent)

    boundaries = []
    for t in range(1, 5):
        # Toy guards ell >= 2^t; the actual t=1000 tower is never constructed.
        first_n = (1 << (1 << (1 << (t - 1)))) - 1
        samples = []
        for n in (first_n - 1, first_n, first_n + 1):
            p = parameters(n)
            passed = p['ell'] >= (1 << t)
            check(f'toy-boundary:{t}:{n}', passed == (n >= first_n))
            samples.append(dict(**p, guard_passed=passed))
        boundaries.append(dict(toy_exponent=t, samples=samples))
    real_sizes = []
    for n in (0, 2, 1000, 10**6, 10**12, (1 << 64)-1, 10**100):
        p = parameters(n)
        p['ell_guard_passed'] = p['ell'] >= 1 << 1000
        p['epsilon_one_tenth_guard_passed'] = p['ell'] * Fraction(1, 10) >= 10
        p['physical_large_input'] = 'false: necessary ell guard fails'
        check(f'practical-ell-guard:{n}', not p['ell_guard_passed'])
        real_sizes.append(p)

    semantics = []
    for label, a, b in [('accent', 'é', 'e'), ('canonical-equivalence', 'é', 'e\u0301'),
                         ('case', 'Depot', 'depot'), ('cutoff-sentinel', 'a'*32, 'b'*32)]:
        record = dict(case=label, source=a, target=b,
                      codepoint_distance=Levenshtein.distance(a, b),
                      utf8_byte_distance=Levenshtein.distance(a.encode(), b.encode()),
                      nfc_distance=Levenshtein.distance(unicodedata.normalize('NFC', a), unicodedata.normalize('NFC', b)),
                      cutoff_3_return=Levenshtein.distance(a, b, score_cutoff=3))
        check(f'semantics:{label}', record['codepoint_distance'] == dp(a, b)
              and record['utf8_byte_distance'] == dp(a.encode(), b.encode()))
        semantics.append(record)
    check('sentinel-is-not-distance', semantics[-1]['cutoff_3_return'] == 4 and semantics[-1]['codepoint_distance'] == 32)
    check('normalization-changes-problem', semantics[1]['codepoint_distance'] == 2 and semantics[1]['nfc_distance'] == 0)

    cases = []
    for n in (4096, 16384, 65536):
        for kind in ('equal', 'sparse_substitutions', 'shifted_periodic', 'all_substitutions'):
            a = 'ab' * (n // 2)
            if kind == 'equal':
                b, known = a, 0
            elif kind == 'sparse_substitutions':
                values = list(a)
                for i in range(0, n, 1024): values[i] = 'c'
                b, known = ''.join(values), (n + 1023)//1024
            elif kind == 'shifted_periodic':
                b, known = 'ba' * (n // 2), 2
            else:
                b, known = 'c' * n, n
            # Each known value has an explicit script and matching symbol-count or
            # elementary lower bound; no general long-string DP oracle is claimed.
            samples = []
            for _ in range(3):
                start = time.perf_counter_ns()
                value = Levenshtein.distance(a, b)
                samples.append(time.perf_counter_ns() - start)
            cutoff_start = time.perf_counter_ns()
            capped = Levenshtein.distance(a, b, score_cutoff=16)
            cutoff_ns = time.perf_counter_ns() - cutoff_start
            check(f'long-generated:{n}:{kind}', value == known and capped == (known if known <= 16 else 17))
            paths = []
            for side, contents in [('source', a), ('target', b)]:
                path = OUT / f'{n}-{kind}-{side}.txt'
                raw = contents.encode('utf-8')
                with path.open('xb') as stream: stream.write(raw)
                paths.append(dict(path=path.name, bytes=len(raw), sha256=hashlib.sha256(raw).hexdigest()))
            cases.append(dict(case=kind, length_each=n, distance=value, cutoff_16_return=capped,
                              exact_ns_samples=samples, exact_ns_median=int(statistics.median(samples)),
                              cutoff_ns=cutoff_ns, input_files=paths))
    source_root = ROOT / 'source/openai-math'
    sources = [
        'preprints/An-Almost-Linear-Approximation-Scheme-for-Edit-Distance-September-24-2026/paper.pdf',
        'lean/docs/121.md', 'lean/ComparatorChallenges/EditApproximation.json',
        'lean/ComparatorChallenges/EditApproximation.lean',
        'lean/OAI/Combinatorics/EditApproximation/Definitions.lean',
        'lean/OAI/Combinatorics/EditApproximation/Geometry/AccuracyCutoff.lean',
        'lean/OAI/Combinatorics/EditApproximation/Specification/RawBandMain.lean']
    payload = dict(generated_at_utc=datetime.now(timezone.utc).isoformat(),
                   runtime=dict(interpreter=sys.executable, python=platform.python_version(),
                                machine=platform.machine(), platform=platform.platform(), rapidfuzz=rapidfuzz.__version__),
                   scope='Existing exact library and selected integer guards; no source approximation execution or accepted source proof.',
                   binary_pairs=len(strings)**2, practical_sizes=real_sizes, toy_boundaries=boundaries,
                   actual_first_ell_guard_length='2^(2^(2^999))-1; necessary ell guard only, not full physical eligibility',
                   accuracy_cutoff='2^(2^max(2^1000,ceil(10/epsilon))); sufficient for two numeric guards only',
                   source_files=[dict(path=p, sha256=hashlib.sha256((source_root/p).read_bytes()).hexdigest()) for p in sources],
                   semantics=semantics, long_generated_cases=cases,
                   benchmark_limit='Three sequential same-process samples on generated strings. No warmup protocol, isolation, source-backend comparison or representative customer workload.')
    jsave(OUT/'comparison.json', payload)
    jsave(REPORT, dict(checks=checks, passed=sum(c['passed'] for c in checks), scope=payload['scope'],
                      comparison=str((OUT/'comparison.json').relative_to(ROOT))))
    print(json.dumps(dict(checks=len(checks), passed=sum(c['passed'] for c in checks), comparison=str(OUT/'comparison.json'))))


if __name__ == '__main__': main()
