"""Collect real local-model outputs; replay deterministic diagnostics without gold labels.

No requests happen at import or during replay. Prompts describe fictional public
scenarios; generated raw strings are retained, including malformed responses.
"""

from __future__ import annotations

import argparse
import hashlib
import json
import math
import platform
import statistics
import sys
import time
import urllib.error
import urllib.request
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "src"))

from narrative_contracts import (  # noqa: E402
    Claim,
    Context,
    DeclaredClaimsConsistent,
    Document,
    ForbiddenPattern,
    LexicalRestatement,
    MinimumTokens,
    Surface,
    evaluate,
)
from narrative_contracts.model import digest  # noqa: E402

PROFILE = "natural-pilot-v1"
SYSTEM = (
    "You write fictional user-facing descriptions of a state machine. The application, not you, "
    "owns all states and transitions. Use only supplied facts. Return one JSON object, no markdown. "
    "Never decide new effects. Use natural informative prose, no boilerplate reassurances."
)


def scenarios():
    """Sixteen paired scenarios, two domains/languages, four contexts per stratum."""
    rows = []
    for domain in ("support", "fiction"):
        for language in ("en", "es"):
            for index in range(4):
                credits = 20 + index * 15
                if domain == "support":
                    setup = (
                        "A fictional store has already authorized two alternatives for a damaged "
                        "parcel: refund store credits or dispatch a replacement. The customer must "
                        "choose one. Refunds become pending, not already received. No delivery date "
                        "has been promised."
                    )
                    options = [
                        {
                            "id": "refund",
                            "action": "Request the authorized store-credit refund",
                            "after": {"status": "refund_pending", "credits": credits + 10},
                        },
                        {
                            "id": "replace",
                            "action": "Request a replacement parcel",
                            "after": {"status": "replacement_pending", "credits": credits},
                        },
                    ]
                else:
                    setup = (
                        "A fictional courier has reached a closed bridge. Two alternatives are "
                        "authorized: pay for the ferry or return to the inn. The package remains "
                        "undelivered in both cases. No character has promised another bridge."
                    )
                    options = [
                        {
                            "id": "ferry",
                            "action": "Pay for the ferry crossing",
                            "after": {"status": "ferry_booked", "credits": credits - 10},
                        },
                        {
                            "id": "inn",
                            "action": "Return to the inn with the package",
                            "after": {"status": "at_inn", "credits": credits},
                        },
                    ]
                scenario = {
                    "id": f"{domain}-{language}-{index}",
                    # Languages/models are paired derivatives of the same scenario family.
                    "group": f"{domain}-{index}",
                    "domain": domain,
                    "language": language,
                    "setup": setup,
                    "current": {"status": "awaiting_choice", "credits": credits},
                    "options": options,
                }
                rows.append(scenario)
    return rows


def prompt(scenario):
    return (
        f"Write all prose in {'Spanish' if scenario['language'] == 'es' else 'English'}. "
        "Describe the scene in 30-50 words and each hypothetical outcome in 25-45 words. "
        "Keep the options mutually exclusive. Add a brief option label. Reflect each supplied "
        "after-state faithfully. State credits as store/game units, not real money. "
        "For every outcome, declare exactly status (string) and credits (integer) in claims; "
        "claims describe that outcome's after-state. Preserve option IDs and their order. "
        'Required JSON shape: {"body": "...", "options": [{"id": "...", "label": "...", '
        '"outcome": "...", "claims": {"status": "...", "credits": 0}}]}.\n'
        + json.dumps(scenario, ensure_ascii=False, sort_keys=True)
    )


def rules():
    outcome_ids = ("outcome.0", "outcome.1")
    return (
        DeclaredClaimsConsistent("claims", outcome_ids),
        MinimumTokens("content", ("body", *outcome_ids), minimum=12, minimum_unique=6),
        ForbiddenPattern(
            "formula",
            (
                r"\bcomo era de esperar\b",
                r"\bsin duda\b",
                r"\bas expected\b",
                r"\bwithout a doubt\b",
            ),
            ("body", *outcome_ids),
        ),
        LexicalRestatement("restatement.0", "label.0", "outcome.0"),
        LexicalRestatement("restatement.1", "label.1", "outcome.1"),
    )


def _pairs(items):
    result = {}
    for key, value in items:
        if key in result:
            raise ValueError(f"duplicate key: {key}")
        result[key] = value
    return result


def _bad_constant(value):
    raise ValueError(f"nonfinite JSON: {value}")


def strict_json(raw):
    return json.loads(raw, object_pairs_hook=_pairs, parse_constant=_bad_constant)


def document(raw, scenario):
    data = strict_json(raw)
    if not isinstance(data, dict) or set(data) != {"body", "options"}:
        raise ValueError("Expected body and options")
    if not isinstance(data["body"], str) or not isinstance(data["options"], list):
        raise ValueError("Invalid scene shape")
    if len(data["options"]) != 2:
        raise ValueError("Expected two branches")
    surfaces = [Surface("body", data["body"])]
    for index, (option, expected) in enumerate(
        zip(data["options"], scenario["options"], strict=True)
    ):
        if not isinstance(option, dict) or set(option) != {"id", "label", "outcome", "claims"}:
            raise ValueError("Invalid option shape")
        if option["id"] != expected["id"]:
            raise ValueError("Branch identity/order mismatch")
        if not isinstance(option["label"], str) or not isinstance(option["outcome"], str):
            raise ValueError("Text must be a string")
        claims = option["claims"]
        if (
            not isinstance(claims, dict)
            or set(claims) != {"status", "credits"}
            or type(claims["status"]) is not str
            or type(claims["credits"]) is not int
        ):
            raise ValueError("Expected a status string and integer credits declaration")
        surfaces.extend(
            (
                Surface(f"label.{index}", option["label"]),
                Surface(
                    f"outcome.{index}",
                    option["outcome"],
                    expected["id"],
                    tuple(Claim(k, v) for k, v in sorted(claims.items())),
                ),
            )
        )
    return Document(tuple(surfaces))


def context(scenario):
    return Context(
        {"current": scenario["current"], **{o["id"]: o["after"] for o in scenario["options"]}}
    )


def _http(base, route, payload=None, timeout=120):
    request = urllib.request.Request(
        base.rstrip("/") + route,
        data=json.dumps(payload).encode() if payload is not None else None,
        headers={"Content-Type": "application/json"},
    )
    with urllib.request.urlopen(request, timeout=timeout) as response:
        return strict_json(response.read())


def _write(path, value):
    path.write_text(json.dumps(value, ensure_ascii=False, indent=2, allow_nan=False) + "\n")


def collect(args):
    if args.output.exists():
        raise FileExistsError(
            "Collection output must be new; failed attempts are never overwritten"
        )
    if len(set(args.models)) != len(args.models):
        raise ValueError("Models must be distinct")
    tags = _http(args.base_url, "/api/tags")
    available = {m["name"]: m for m in tags["models"]}
    missing = set(args.models) - available.keys()
    if missing:
        raise ValueError(f"Pull these models first: {sorted(missing)}")
    # Collect full public model metadata to pin templates/quantization and weight digests.
    selected = {
        name: {"tag": available[name], "show": _http(args.base_url, "/api/show", {"model": name})}
        for name in args.models
    }
    dataset = scenarios()
    manifest = {
        "schema_version": 1,
        "profile": PROFILE,
        "collection_started_utc": datetime.now(timezone.utc).isoformat(),
        "ollama_version": _http(args.base_url, "/api/version"),
        "python": platform.python_version(),
        "platform": platform.platform(),
        "models": selected,
        "scenarios": dataset,
        "scenario_digest": digest(dataset),
        "system": SYSTEM,
        "seed": args.seed,
        "options": {"seed": args.seed, "temperature": 0, "num_predict": 512, "num_ctx": 4096},
        "format": "json",
        "retries": 0,
        "repairs": 0,
        "protocol": "paper/release-protocol.md",
        "protocol_sha256": hashlib.sha256(
            (ROOT / "paper/release-protocol.md").read_bytes()
        ).hexdigest(),
        "runner_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
        "rule_configuration": [r.configuration() for r in rules()],
        "label_provenance": "No human or model-judge labels; real model outputs on authored scenarios",
        "cost": {"provider_fee_usd": 0, "electricity_and_hardware_cost": None},
    }
    args.output.mkdir(parents=True)
    _write(args.output / "manifest.json", manifest)  # Freeze before the first model output.
    with (args.output / "outputs.jsonl").open("x") as ledger:
        for model in args.models:
            for scenario in dataset:
                request = {
                    "model": model,
                    "system": SYSTEM,
                    "prompt": prompt(scenario),
                    "stream": False,
                    "format": "json",
                    "options": manifest["options"],
                    "keep_alive": "5m",
                }
                capabilities = selected[model]["show"].get("capabilities", [])
                if "thinking" in capabilities:
                    request["think"] = False
                row = {
                    "id": f"{model}/{scenario['id']}",
                    "model": model,
                    "scenario_id": scenario["id"],
                    "group": scenario["group"],
                    "request": request,
                    "request_digest": digest(request),
                }
                started = time.perf_counter()
                try:
                    row["response"] = _http(args.base_url, "/api/generate", request, args.timeout)
                    row["collection_status"] = "returned"
                except (urllib.error.URLError, TimeoutError, ValueError) as exc:
                    row["collection_status"] = "error"
                    row["error_type"] = type(exc).__name__
                row["wall_seconds"] = time.perf_counter() - started
                ledger.write(json.dumps(row, ensure_ascii=False, allow_nan=False) + "\n")
                ledger.flush()
                print(row["id"], row["collection_status"], flush=True)
    replay(args.output, args.output / "evaluation")


def evaluate_record(record, scenario):
    result = {
        "id": record["id"],
        "model": record["model"],
        "scenario_id": scenario["id"],
        "domain": scenario["domain"],
        "language": scenario["language"],
        "group": scenario["group"],
        "status": "collection_error",
    }
    if record.get("collection_status") != "returned":
        return result
    response = record.get("response", {})
    if not response.get("done") or response.get("done_reason") == "length":
        return {**result, "status": "incomplete_generation"}
    try:
        doc = document(response["response"], scenario)
    except (ValueError, TypeError, KeyError) as exc:
        return {**result, "status": "invalid_structure", "error_type": type(exc).__name__}
    ctx = context(scenario)
    configured = rules()
    strict = evaluate(doc, ctx, configured)
    invariant = evaluate(doc, ctx, configured[:1])
    heuristic = evaluate(doc, ctx, configured[1:])
    return {
        **result,
        "status": "evaluated",
        "report": strict.to_dict(),
        "invariant_accepted": invariant.accepted,
        "heuristic_accepted": heuristic.accepted,
        "length_assertions_accepted": len(doc.surface("body").text.strip()) >= 70
        and all(len(doc.surface(f"outcome.{i}").text.strip()) >= 60 for i in range(2)),
    }


def summarize(rows, records):
    result = {}
    for model in sorted({r["model"] for r in rows}):
        selected = [r for r in rows if r["model"] == model]
        valid = [r for r in selected if r["status"] == "evaluated"]
        raw = [r for r in records if r["model"] == model]
        seconds = sorted(r["wall_seconds"] for r in raw)
        findings = Counter(
            c["code"] for row in valid for c in row["report"]["checks"] if c["status"] == "violated"
        )
        result[model] = {
            "attempts": len(selected),
            "status_counts": dict(Counter(r["status"] for r in selected)),
            "structured_outputs": len(valid),
            "strict_accepted": sum(r["report"]["accepted"] for r in valid),
            "invariant_accepted": sum(r["invariant_accepted"] for r in valid),
            "heuristic_accepted": sum(r["heuristic_accepted"] for r in valid),
            "length_assertions_accepted": sum(r["length_assertions_accepted"] for r in valid),
            "violation_findings": dict(sorted(findings.items())),
            "generation_wall_median_s": statistics.median(seconds) if seconds else None,
            "generation_wall_p95_s": seconds[max(0, math.ceil(0.95 * len(seconds)) - 1)]
            if seconds
            else None,
            "output_tokens": sum(r.get("response", {}).get("eval_count", 0) for r in raw),
            "prompt_tokens": sum(r.get("response", {}).get("prompt_eval_count", 0) for r in raw),
        }
    return result


def replay(input_dir, output_dir):
    if output_dir.exists():
        raise FileExistsError("Replay output must be new")
    manifest = strict_json((input_dir / "manifest.json").read_text())
    if digest(manifest["scenarios"]) != manifest["scenario_digest"]:
        raise ValueError("Scenario digest mismatch")
    if manifest["profile"] != PROFILE or manifest["rule_configuration"] != [
        r.configuration() for r in rules()
    ]:
        raise ValueError("Replay requires the frozen rule configuration")
    dataset = {s["id"]: s for s in manifest["scenarios"]}
    records = [strict_json(line) for line in (input_dir / "outputs.jsonl").read_text().splitlines()]
    ids = [r["id"] for r in records]
    if len(ids) != len(set(ids)):
        raise ValueError("Duplicate collection row")
    rows, timings = [], []
    for record in records:
        if record["request_digest"] != digest(record["request"]):
            raise ValueError("Request digest mismatch")
        started = time.perf_counter_ns()
        rows.append(evaluate_record(record, dataset[record["scenario_id"]]))
        timings.append((time.perf_counter_ns() - started) / 1_000_000)
    expected = {f"{m}/{s}" for m in manifest["models"] for s in dataset}
    summary = {
        "profile": PROFILE,
        "corpus_digest": digest(records),
        "scenario_groups": len({s["group"] for s in dataset.values()}),
        "missing_attempts": sorted(expected - set(ids)),
        "models": summarize(rows, records),
        "limitations": [
            "No independent labels: acceptance is not accuracy or quality",
            "Small authored pilot; no production traffic or model ranking",
            "Claims check annotations, not whether prose entails them",
            "Regeneration is hardware/version dependent; replay of frozen records is deterministic",
            "No judge baseline; no superiority or semantic FPR/FNR claim",
        ],
    }
    output_dir.mkdir(parents=True)
    _write(output_dir / "summary.json", summary)
    _write(output_dir / "reports.json", rows)
    _write(
        output_dir / "latency.json",
        {
            "samples": len(timings),
            "median_ms": statistics.median(timings) if timings else None,
            "p95_ms": sorted(timings)[max(0, math.ceil(0.95 * len(timings)) - 1)]
            if timings
            else None,
            "note": "Warm local replay, parsing and three profiles included; timing not deterministic",
        },
    )
    print(json.dumps(summary, indent=2))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    sub = parser.add_subparsers(dest="command", required=True)
    collect_parser = sub.add_parser("collect")
    collect_parser.add_argument("--output", type=Path, required=True)
    collect_parser.add_argument("--base-url", default="http://localhost:11434")
    collect_parser.add_argument("--models", nargs="+", required=True)
    collect_parser.add_argument("--seed", type=int, default=1729)
    collect_parser.add_argument("--timeout", type=float, default=120)
    replay_parser = sub.add_parser("replay")
    replay_parser.add_argument("--input", type=Path, required=True)
    replay_parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()
    if args.command == "collect":
        collect(args)
    else:
        replay(args.input, args.output)


if __name__ == "__main__":
    main()
