"""Phase 4.5 — DS-STAR-style verification pass over the teacher outputs.

A cheap Gemini Flash batch job audits every generated analysis BEFORE the
forward-return filter. It catches what return-based filtering can't:
  - figures quoted in the analysis that don't match the provided fundamentals
  - references to events after the as-of date (lookahead leakage)
  - missing or unparseable VERDICT blocks

Unlike full DS-STAR, this is a single batch pass (50% off pricing), not an
iterative agent loop — verification without the token blowup. Cost: ~$3.

    python src/data/verify_teacher_outputs.py                     # submit + poll + filter
    python src/data/verify_teacher_outputs.py --collect batches/...

Input:  datasets/teacher_outputs.jsonl
Output: datasets/teacher_outputs_verified.jsonl  (passing rows only)
"""
import argparse
import json
import time
from pathlib import Path

from google import genai
from google.genai import types

ROOT = Path(__file__).resolve().parents[2]
IN = ROOT / "datasets" / "teacher_outputs.jsonl"
BATCH_FILE = ROOT / "datasets" / "gemini_verify_requests.jsonl"
OUT = ROOT / "datasets" / "teacher_outputs_verified.jsonl"

MODEL = "gemini-3.5-flash"

VERIFIER_INSTRUCTIONS = """You are auditing an investment analysis for data integrity, \
not for whether its opinion is correct.

The analysis was written as of {as_of} using ONLY the fundamentals data shown below. Check:
1. NUMBERS: every specific financial figure in the analysis is traceable to the provided \
fundamentals (or is clearly derived from them, e.g. margins, ratios). Flag invented figures.
2. LOOKAHEAD: the analysis must not reference events, results, prices, or facts from after \
{as_of}. Flag anything the author could not have known on that date.
3. FORMAT: the analysis ends with a parseable block: VERDICT: {{"direction": ..., \
"conviction": ..., "horizon_months": ...}}

Respond with only JSON: {{"pass": true/false, "reasons": ["..."]}} — reasons empty if pass.

## Fundamentals provided to the author
{fundamentals}

## Analysis to audit
{analysis}"""

COMPLETED = {"JOB_STATE_SUCCEEDED", "JOB_STATE_FAILED",
             "JOB_STATE_CANCELLED", "JOB_STATE_EXPIRED"}


def load_rows() -> list[dict]:
    return [json.loads(l) for l in IN.read_text(encoding="utf-8").splitlines()]


def fundamentals_from_user(user: str) -> str:
    start = user.find("{")
    end = user.find("## Price context")
    return user[start:end].rsplit("}", 1)[0] + "}"


def write_batch_file(rows: list[dict]) -> None:
    with BATCH_FILE.open("w", encoding="utf-8") as fh:
        for i, row in enumerate(rows):
            prompt = VERIFIER_INSTRUCTIONS.format(
                as_of=row["as_of"],
                fundamentals=fundamentals_from_user(row["user"]),
                analysis=row["assistant"],
            )
            fh.write(json.dumps({
                "key": f"v{i:05d}",
                "request": {
                    "contents": [{"role": "user", "parts": [{"text": prompt}]}],
                    "generation_config": {
                        "temperature": 0.0,
                        "response_mime_type": "application/json",
                    },
                },
            }) + "\n")


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--collect", metavar="JOB_NAME")
    args = ap.parse_args()

    client = genai.Client()
    rows = load_rows()

    if args.collect:
        job_name = args.collect
    else:
        write_batch_file(rows)
        uploaded = client.files.upload(
            file=str(BATCH_FILE),
            config=types.UploadFileConfig(display_name="deepfeline-verify",
                                          mime_type="jsonl"),
        )
        job = client.batches.create(model=MODEL, src=uploaded.name,
                                    config={"display_name": "deepfeline-verify"})
        job_name = job.name
        print(f"Submitted {len(rows)} audits. Job: {job_name}")

    while True:
        job = client.batches.get(name=job_name)
        if job.state.name in COMPLETED:
            break
        print(f"state: {job.state.name} ...")
        time.sleep(60)

    if job.state.name != "JOB_STATE_SUCCEEDED":
        raise SystemExit(f"Batch ended in {job.state.name}")

    raw = client.files.download(file=job.dest.file_name).decode("utf-8")
    verdicts: dict[int, dict] = {}
    for line in raw.splitlines():
        r = json.loads(line)
        if "response" not in r:
            continue
        try:
            parts = r["response"]["candidates"][0]["content"]["parts"]
            text = "".join(p.get("text", "") for p in parts if not p.get("thought"))
            verdicts[int(r["key"][1:])] = json.loads(text)
        except (KeyError, IndexError, json.JSONDecodeError):
            continue

    kept, dropped = 0, 0
    with OUT.open("w", encoding="utf-8") as fh:
        for i, row in enumerate(rows):
            v = verdicts.get(i)
            if v is None or not v.get("pass"):
                dropped += 1
                if v and v.get("reasons"):
                    print(f"{row['ticker']} {row['as_of']}: {v['reasons'][0]}")
                continue
            fh.write(json.dumps(row) + "\n")
            kept += 1

    print(f"\nkept {kept}, dropped {dropped} -> {OUT}")


if __name__ == "__main__":
    main()
