#!/usr/bin/env python3
"""Fresh Gemini 3.1 Pro forward/reverse replay; retained other models untouched.

Native Gemini REST, high thinking, no retries/fallback, four concurrent calls.
Only experiment artifacts are written; production database access is read-only.
"""
from __future__ import annotations

import argparse
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import datetime, timezone
from pathlib import Path

import httpx
from dotenv import load_dotenv
from jsonschema import validate

BACKEND = Path(__file__).resolve().parents[2]
ROOT = BACKEND.parent
sys.path.insert(0, str(BACKEND))
sys.path.insert(0, str(Path(__file__).resolve().parent))

from astra_judge_pilot import compare, digest, manifest_parts, save

MODEL = "gemini-3.1-pro-preview"
CONFIG = {"temperature": 1.0, "maxOutputTokens": 16000,
          "thinkingConfig": {"thinkingLevel": "high"},
          "responseMimeType": "application/json"}


def native_parts(parts):
    converted = []
    for part in parts:
        if part["type"] == "text":
            converted.append({"text": part["text"]})
        else:
            header, data = part["image_url"]["url"].split(",", 1)
            converted.append({"inlineData": {"mimeType": header[5:].split(";")[0], "data": data}})
    return converted


def ask(job, system, schema, single=False, max_tokens=None):
    from app.modules.llm.gemini_key_pool import get_next_key
    from app.modules.pipeline.multiroll_select import normalize_flip_verdict

    tag, order, parts, identity = job
    mapping = {} if single else ({"A": "A", "B": "B"} if order == "forward" else {"A": "B", "B": "A"})
    key = get_next_key()
    slot = {"model": MODEL, "reasoning_effort": "high", "order": order,
            "display_to_canonical": mapping, "request_identity": identity,
            "started_at": datetime.now(timezone.utc).isoformat(), "ok": False,
            "http_attempt_count": 0}
    started = time.monotonic()
    try:
        body = {"systemInstruction": {"parts": [{"text": system}]},
                "contents": [{"role": "user", "parts": native_parts(parts)}],
                "generationConfig": {**CONFIG, **({"maxOutputTokens": max_tokens} if max_tokens else {}),
                                     "responseJsonSchema": schema}}
        slot["request_options"] = {k: v for k, v in body["generationConfig"].items() if k != "responseJsonSchema"}
        with httpx.Client(timeout=600, transport=httpx.HTTPTransport(retries=0)) as client:
            slot["http_attempt_count"] += 1
            response = client.post(
                f"https://generativelanguage.googleapis.com/v1beta/models/{MODEL}:generateContent",
                headers={"x-goog-api-key": key, "Content-Type": "application/json"}, json=body)
        slot["http_status"] = response.status_code
        payload = response.json()
        slot["response"] = payload
        response.raise_for_status()
        slot["returned_model"] = payload.get("modelVersion")
        slot["response_id"] = payload.get("responseId")
        native_usage = payload.get("usageMetadata", {})
        thoughts = native_usage.get("thoughtsTokenCount", 0)
        completion = native_usage.get("candidatesTokenCount", 0)
        slot["usage"] = {"prompt_tokens": native_usage.get("promptTokenCount", 0),
                         "completion_tokens": completion + thoughts,
                         "completion_tokens_details": {"reasoning_tokens": thoughts},
                         "native": native_usage}
        candidate = payload.get("candidates", [{}])[0]
        slot["finish_reason"] = candidate.get("finishReason")
        if slot["finish_reason"] != "STOP":
            raise ValueError(f"Incomplete response: {slot['finish_reason']}")
        answer = "".join(p.get("text", "") for p in candidate.get("content", {}).get("parts", [])
                         if not p.get("thought"))
        verdict = json.loads(answer)
        validate(verdict, schema)
        slot.update(raw=verdict, normalized=verdict if single else normalize_flip_verdict(verdict, mapping, ["A", "B"]), ok=True)
    except Exception as exc:
        slot.update(error_type=type(exc).__name__, error=str(exc).replace(key, "[REDACTED]"))
    slot["elapsed_seconds"] = round(time.monotonic() - started, 3)
    return tag, order, slot


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--baseline", type=Path, default=ROOT / "artifact/20260905_astra_xhigh_vlm_compare")
    parser.add_argument("--out", type=Path, default=ROOT / "artifact/20260905_vlm_four_model_compare")
    parser.add_argument("--workers", type=int, default=4)
    parser.add_argument("--live", action="store_true")
    args = parser.parse_args()
    if not 1 <= args.workers <= 8 or args.out.resolve() == args.baseline.resolve():
        parser.error("workers must be 1..8 and output must differ from baseline")
    load_dotenv(BACKEND / ".env", override=False)
    from muse_judge_pilot import _parts_for

    original_manifest = json.loads((args.baseline / "input_manifest.json").read_text())
    baseline_results = json.loads((args.baseline / "results.json").read_text())
    records_path = Path(original_manifest["records"])
    records = json.loads(records_path.read_text())
    system = (args.baseline / "judge_sys.txt").read_text()
    schema = json.loads((args.baseline / "judge_schema.json").read_text())
    headers = json.loads((args.baseline / "headers.json").read_text())
    sources = {str(p): digest(p.read_bytes()) for p in
               [records_path, args.baseline / "results.json", args.baseline / "input_manifest.json"]}
    args.out.mkdir(parents=True, exist_ok=True)
    for name in ("judge_sys.txt", "judge_schema.json", "headers.json"):
        (args.out / name).write_bytes((args.baseline / name).read_bytes())
    path = args.out / "results.json"
    results = json.loads(path.read_text()) if path.exists() else baseline_results
    mine = results.setdefault(MODEL, {})
    jobs, inputs = [], {}
    for tag in original_manifest["shots"]:
        for order in ("forward", "reverse"):
            parts = _parts_for(records[tag], records_path.parent, tag, ["A", "B"],
                               ["A", "B"] if order == "forward" else ["B", "A"],
                               headers["default"], headers["bgfirst"])
            retained = manifest_parts(parts)
            if retained != original_manifest["inputs"][f"{tag}_{order}"]["parts"]:
                raise ValueError(f"{tag} {order}: text/image input differs from Astra experiment")
            identity = digest({"model": MODEL, "config": CONFIG, "schema": schema,
                               "system": system, "parts": retained})
            inputs[f"{tag}_{order}"] = {"request_identity": identity, "parts": retained}
            old = mine.get(tag, {}).get(order, {})
            if old.get("ok"):
                if old.get("request_identity") != identity:
                    raise ValueError("Input changed: choose a new output directory")
                continue
            jobs.append((tag, order, parts, identity))
    if len(jobs) > 12:
        raise ValueError("More than 12 calls are outside this replay scope")
    manifest = {**original_manifest, "model": MODEL, "effort": "high", "api": "generateContent",
                "base_url": "https://generativelanguage.googleapis.com/v1beta",
                "generationConfig": CONFIG, "sdk_max_retries": 0,
                "pending_calls": len(jobs), "workers": args.workers, "inputs": inputs,
                "source_manifest": str(args.baseline / "input_manifest.json"),
                "historical_source_sha256": sources, "same_inputs_as_astra": True}
    for obsolete in ("key_slot", "max_completion_tokens"):
        manifest.pop(obsolete, None)
    save(args.out / "input_manifest.json", manifest)
    print(json.dumps({"model": MODEL, "thinking": "high", "calls": len(jobs),
                      "workers": args.workers, "same_inputs_as_astra": True, "live": args.live}), flush=True)
    if not args.live:
        return 0
    started = time.monotonic()
    with ThreadPoolExecutor(max_workers=args.workers) as pool:
        for future in as_completed([pool.submit(ask, j, system, schema) for j in jobs]):
            tag, order, slot = future.result()
            mine.setdefault(tag, {})[order] = slot
            save(path, results)
            print(json.dumps({"tag": tag, "order": order, "ok": slot["ok"],
                  "winner": slot.get("normalized", {}).get("winner"), "seconds": slot["elapsed_seconds"],
                  "http_status": slot.get("http_status"), "error": slot.get("error")}, ensure_ascii=False), flush=True)
    summary = compare(results, records, original_manifest["shots"])
    summary.update(this_run_wall_seconds=round(time.monotonic()-started, 3), this_run_call_count=len(jobs),
                   historical_sources_unchanged=all(digest(Path(p).read_bytes()) == h for p, h in sources.items()))
    save(args.out / "summary.json", summary)
    print(json.dumps(summary["models"], ensure_ascii=False, indent=2), flush=True)
    return 0 if all(mine.get(t, {}).get(o, {}).get("ok") for t in original_manifest["shots"]
                    for o in ("forward", "reverse")) else 1


if __name__ == "__main__":
    raise SystemExit(main())
