"""Small, offline-after-model-download teaching labs. Never queries AI platforms."""
import argparse
from collections import defaultdict
from datetime import datetime, timezone
import hashlib
import importlib.metadata
import json
import math
from pathlib import Path
import platform
import re
import sqlite3
import time

HERE = Path(__file__).resolve().parent


def load_data():
    return json.loads((HERE / 'data.json').read_text())


def ranking_metrics(ranked, grades, k=3):
    if k < 1 or len(ranked) != len(set(ranked)):
        raise ValueError('Use a positive k and unique document IDs')
    relevant = {key for key, grade in grades.items() if grade > 0}
    if not relevant:
        return dict(recall=None, mrr=None, ndcg=None)
    top = ranked[:k]
    dcg = sum((2 ** grades.get(key, 0) - 1) / math.log2(i + 2) for i, key in enumerate(top))
    ideal = sum((2 ** g - 1) / math.log2(i + 2) for i, g in enumerate(sorted(grades.values(), reverse=True)[:k]))
    return dict(recall=len(set(top) & relevant) / len(relevant),
                mrr=next((1 / (i + 1) for i, key in enumerate(top) if key in relevant), 0),
                ndcg=dcg / ideal)


def rrf(rankings, c=60):
    if c <= 0:
        raise ValueError('Fusion constant must be positive')
    scores = defaultdict(float)
    for ranking in rankings:
        seen = set()
        for rank, key in enumerate(ranking, 1):
            if key in seen:
                continue
            seen.add(key)
            scores[key] += 1 / (c + rank)
    return sorted(scores, key=lambda key: (-scores[key], key))


def lexical_rank(documents, query, exact_model=False):
    terms = re.findall(r'[A-Za-z0-9]+', query)
    if not terms:
        return []
    with sqlite3.connect(':memory:') as db:
        db.execute('CREATE VIRTUAL TABLE products USING fts5(id UNINDEXED, model, body)')
        db.executemany('INSERT INTO products VALUES (?,?,?)', [(d['id'], d['model'], d['text']) for d in documents])
        match = ' OR '.join('"' + t + '"' for t in terms)
        ranks = [row[0] for row in db.execute('SELECT id FROM products WHERE products MATCH ? ORDER BY bm25(products,0,4,1),id', (match,))]
    # Explicit identifier equality is a business constraint, not a BM25 capability.
    if exact_model:
        exact = [d['id'] for d in documents if d['model'].casefold() == query.casefold().strip()]
        ranks = exact + [key for key in ranks if key not in exact]
    return ranks


def wilson(x, n, z=1.959963984540054):
    if type(x) is not int or type(n) is not int or n <= 0 or x < 0 or x > n:
        raise ValueError('Expected 0 <= successes <= positive denominator')
    p = x / n
    denominator = 1 + z * z / n
    center = (p + z * z / (2 * n)) / denominator
    delta = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / denominator
    return [max(0, center - delta), min(1, center + delta)]


def bilingual_errors(zh, en):
    fields = ['entity_id', 'model', 'voltage_v', 'tank_l', 'scope', 'source_version', 'evidence_id']
    if any(key not in zh or key not in en for key in fields):
        raise ValueError('Missing invariant fact field')
    return [key for key in fields if zh[key] != en[key]]


def governance_lab(data):
    import jsonschema
    schema = json.loads((HERE / 'fact.schema.json').read_text())
    validator = jsonschema.Draft202012Validator(schema)
    for fact in data['facts'].values():
        validator.validate(fact)
    sources = {record['id'] for record in data['documents']}
    claim_ids = set()
    for claim in data['claims']:
        if (not claim.get('id') or claim['id'] in claim_ids or not claim.get('text')
                or not claim.get('reason') or not isinstance(claim.get('sources'), list)
                or not claim['sources'] or not set(claim['sources']) <= sources
                or claim.get('human_label') not in {'supported', 'contradicted', 'insufficient', 'partly-supported'}):
            raise ValueError('Invalid claim annotation or unknown source reference')
        claim_ids.add(claim['id'])
    invalid = {**data['facts']['en'], 'voltage_v': '220V'}
    stale = {**data['facts']['en'], 'source_version': 'v0'}
    return dict(invalid_fact_rejected=bool(list(validator.iter_errors(invalid))),
                stale_translation_rejected=bool(bilingual_errors(data['facts']['zh'], stale)),
                correct_translation_passes=not bilingual_errors(**data['facts']),
                claim_labels=data['claims'],
                limitation='Claim labels are author-reviewed annotations, not automated entailment or truth verification.')


def measurement_lab(data):
    import numpy as np
    runs = data['synthetic_runs']
    before = [value for row in runs for value in row['before']]
    after = [value for row in runs for value in row['after']]
    delta_by_question = np.array([np.mean(row['after']) - np.mean(row['before']) for row in runs])
    rng = np.random.default_rng(20261002)
    samples = rng.choice(delta_by_question, (10000, len(runs)), replace=True).mean(axis=1)
    return dict(questions=len(runs), runs_per_period=len(before),
                before_rate=float(np.mean(before)), after_rate=float(np.mean(after)),
                paired_change=float(delta_by_question.mean()),
                paired_question_bootstrap_95=np.quantile(samples, [.025, .975]).tolist(),
                bootstrap_seed=20261002, bootstrap_replicates=10000,
                independent_example={'successes': 2, 'n': 10, 'wilson_95': wilson(2, 10)},
                limitation='Synthetic repeated records. Resample paired questions, not individual repeated runs. Six clusters are too few for reliable real-world inference. No causal claim.')


def chunking_lab(data):
    text = data['documents'][0]['text']
    fixed = [text[i:i+64] for i in range(0, len(text), 64)]
    structured = ['AX-220 | electrical | Supply 220 V only. | D1 | v1',
                  'AX-220 | use | Indoor sealed hard floors. Do not use outdoors or on untreated wood. | D1 | v1',
                  'AX-220 | tank | Tank 20 L. | D1 | v1']
    conditions = {'voltage': ['220 V only'], 'capacity': ['20 L'],
                  'usage_limits': ['Indoor sealed hard floors', 'Do not use outdoors or on untreated wood']}
    def checks(chunks):
        return {key: any('AX-220' in chunk and all(term in chunk for term in terms) for chunk in chunks)
                for key, terms in conditions.items()}
    comparisons = []
    for window in [32, 64, 128]:
        chunks = [text[i:i+window] for i in range(0, len(text), window)]
        comparisons.append(dict(window=window, chunks=chunks, standalone_checks=checks(chunks)))
    return dict(fixed_character_window=64, fixed=fixed, structured=structured,
                comparisons=comparisons, structured_checks=checks(structured),
                note='Standalone checks require the model and full, author-selected phrases in one chunk. This literal check is not semantic entailment or retrieval/generation evaluation. Structured chunks are explicitly authored, not an automatic document parser. No universal web paragraph length is implied.')


def retrieval_lab(data, offline=False):
    import numpy as np
    import torch
    from huggingface_hub import HfApi
    from sentence_transformers import SentenceTransformer, CrossEncoder
    torch.manual_seed(20261002)
    torch.set_num_threads(2)
    pins_file = HERE / 'model-revisions.json'
    if pins_file.exists():
        models = json.loads(pins_file.read_text())
    elif offline:
        raise ValueError('Model revision manifest missing; run online once before offline reproduction')
    else:
        models = {key: {'name': name, 'revision': HfApi().model_info(name).sha} for key, name in {
            'embedding': 'sentence-transformers/all-MiniLM-L6-v2',
            'reranker': 'cross-encoder/ms-marco-MiniLM-L6-v2'}.items()}
        pins_file.write_text(json.dumps(models, indent=2) + '\n')
    encoder = SentenceTransformer(models['embedding']['name'], revision=models['embedding']['revision'], device='cpu', local_files_only=offline)
    reranker = CrossEncoder(models['reranker']['name'], revision=models['reranker']['revision'], device='cpu', local_files_only=offline)
    docs = data['documents']
    ids = [d['id'] for d in docs]
    start = time.perf_counter()
    vectors = encoder.encode([d['text'] for d in docs], normalize_embeddings=True)
    corpus_encoding_ms = (time.perf_counter() - start) * 1000
    result = []
    for query in data['queries']:
        start = time.perf_counter()
        lexical = lexical_rank(docs, query['text'])
        lexical_ms = (time.perf_counter() - start) * 1000
        start = time.perf_counter()
        vector = encoder.encode(query['text'], normalize_embeddings=True)
        scores = vectors @ vector
        dense = [ids[i] for i in sorted(range(len(ids)), key=lambda i: (-float(scores[i]), ids[i]))]
        dense_ms = (time.perf_counter() - start) * 1000
        start = time.perf_counter()
        fused = rrf([lexical[:5], dense[:5]])
        fusion_ms = (time.perf_counter() - start) * 1000
        candidates = fused[:5]
        start = time.perf_counter()
        cross = reranker.predict([(query['text'], docs[ids.index(key)]['text']) for key in candidates])
        reordered = [candidates[i] for i in sorted(range(len(candidates)), key=lambda i: (-float(cross[i]), candidates[i]))]
        cross_ms = (time.perf_counter() - start) * 1000
        rankings = dict(bm25=lexical, dense=dense, rrf=fused, reranked=reordered)
        result.append(dict(query_id=query['id'], group=query['group'], text=query['text'], **rankings,
            dense_scores={key:float(scores[ids.index(key)]) for key in dense},
            cross_encoder_scores={key:float(cross[candidates.index(key)]) for key in candidates},
            metrics_at_3={name:ranking_metrics(ranking, query['grades']) for name, ranking in rankings.items()},
            latency_ms=dict(bm25_with_index=lexical_ms, dense_query=dense_ms, rrf=fusion_ms, rerank=cross_ms),
            unjudged_no_answer=not query['grades']))
    return models, result, corpus_encoding_ms


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--offline', action='store_true', help='Use cached, pinned models only')
    parser.add_argument('--output', type=Path, default=HERE / 'results.json')
    args = parser.parse_args()
    data = load_data()
    models, retrieval, encoding_ms = retrieval_lab(data, args.offline)
    out = dict(data_kind='synthetic-teaching', external_ai_tested=False,
               created_at=datetime.now(timezone.utc).isoformat(),
               inputs_sha256={name:hashlib.sha256((HERE / name).read_bytes()).hexdigest() for name in ['data.json','fact.schema.json','labs.py','model-revisions.json']},
               environment={'python':platform.python_version(), 'platform':platform.platform(), 'sqlite':sqlite3.sqlite_version,
                            'packages':{p:importlib.metadata.version(p) for p in ['torch','numpy','sentence-transformers','transformers','jsonschema','huggingface-hub']}},
               models=models, retrieval=retrieval, corpus_encoding_ms=encoding_ms,
               gates=governance_lab(data), statistics=measurement_lab(data), chunking=chunking_lab(data),
               limitations=['Nine synthetic English records; ten authored queries, including one Chinese stress query and one no-answer query.',
                            'No-answer query is excluded from ranking averages, never treated as perfect recall.',
                            'Author labels are not independent annotator consensus. No held-out tuning or external validity.',
                            'CPU timings are single-run observations. Index construction is included only for lexical timing; no speed comparison claim.',
                            'English embedding/reranking models are not validated as Chinese search engines. A nearest neighbour is not a factual answer.'])
    args.output.write_text(json.dumps(out, ensure_ascii=False, indent=2) + '\n')
    print(json.dumps({'queries':len(retrieval), 'models':models, 'gates':{k:v for k,v in out['gates'].items() if isinstance(v,bool)}}, indent=2))


if __name__ == '__main__':
    main()
