"""Phase 4 — generate teacher analyses with Gemini via the Batch API.

Runs LOCALLY (no GPU). Requires GEMINI_API_KEY in the environment
(create one at https://aistudio.google.com/apikey — needs a paid/billed key;
the free tier's daily limits won't cover a full run).

Batch mode runs at 50% of standard token prices with up to 24h turnaround:
Gemini 3.1 Pro batched ~= $1 / $6 per million tokens.

    python src/train/generate_teacher_data.py                    # submit + poll + collect
    python src/train/generate_teacher_data.py --collect batches/...   # resume

Input:  datasets/teacher_prompts.jsonl  (from Phase 3)
Output: datasets/teacher_outputs.jsonl  ({ticker, as_of, system, user, assistant})

Each prompt is sampled len(LENSES) times through different analytical lenses so
Phase 5's rejection sampling has differentiated candidates to filter.
"""
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]
PROMPTS = ROOT / "datasets" / "teacher_prompts.jsonl"
BATCH_FILE = ROOT / "datasets" / "gemini_batch_requests.jsonl"
OUT = ROOT / "datasets" / "teacher_outputs.jsonl"

MODEL = "gemini-3.1-pro-preview"

LENSES = [
    "Weight your standard process normally.",
    "Lead with the bear case: assume the market is right to be pessimistic and "
    "make the thesis survive that assumption before turning constructive.",
    "Lead with asymmetry: focus on what the payoff distribution looks like if "
    "the consensus narrative is wrong in either direction.",
]

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


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


def write_batch_file(prompts: list[dict]) -> None:
    with BATCH_FILE.open("w", encoding="utf-8") as fh:
        for i, p in enumerate(prompts):
            for k, lens in enumerate(LENSES):
                fh.write(json.dumps({
                    "key": f"p{i:05d}_s{k}",
                    "request": {
                        "system_instruction": {"parts": [{"text": p["system"]}]},
                        "contents": [{
                            "role": "user",
                            "parts": [{"text": f"{p['user']}\n\n"
                                               f"Analytical lens for this pass: {lens}"}],
                        }],
                        "generation_config": {"temperature": 0.8},
                    },
                }) + "\n")


def extract_text(response: dict) -> str:
    parts = response["candidates"][0]["content"]["parts"]
    return "".join(part.get("text", "") for part in parts
                   if not part.get("thought"))


def collect(client: genai.Client, job, prompts: list[dict]) -> int:
    raw = client.files.download(file=job.dest.file_name).decode("utf-8")
    written = 0
    with OUT.open("a", encoding="utf-8") as fh:
        for line in raw.splitlines():
            row = json.loads(line)
            if "error" in row or "response" not in row:
                print(f"{row.get('key')}: {row.get('error', 'no response')}")
                continue
            try:
                text = extract_text(row["response"])
            except (KeyError, IndexError):
                print(f"{row.get('key')}: unexpected response shape, skipping")
                continue
            if not text.strip():
                continue
            p = prompts[int(row["key"][1:6])]
            fh.write(json.dumps({
                "ticker": p["ticker"],
                "as_of": p["as_of"],
                "system": p["system"],
                "user": p["user"],
                "assistant": text,
            }) + "\n")
            written += 1
    return written


def run_batch(client: genai.Client, chunk: list[dict], label: str) -> int:
    """Submit one batch for `chunk`, poll to completion, collect. Returns rows written."""
    write_batch_file(chunk)
    uploaded = client.files.upload(
        file=str(BATCH_FILE),
        config=types.UploadFileConfig(display_name=label, mime_type="jsonl"),
    )
    job = None
    for attempt in range(6):  # 429s clear as earlier batches drain from the queue
        try:
            job = client.batches.create(model=MODEL, src=uploaded.name,
                                        config={"display_name": label})
            break
        except Exception as e:  # noqa: BLE001
            if "429" in str(e) or "RESOURCE_EXHAUSTED" in str(e):
                print(f"{label}: quota busy, retrying in 3 min ({attempt + 1}/6)")
                time.sleep(180)
            else:
                raise
    if job is None:
        raise SystemExit(f"{label}: still quota-limited after 6 attempts")

    print(f"{label}: {len(chunk) * len(LENSES)} requests submitted, job {job.name}")
    job_name = job.name
    while True:
        try:
            job = client.batches.get(name=job_name)
        except Exception as e:  # noqa: BLE001 — transient network blips must not kill the run
            print(f"{label}: poll failed ({type(e).__name__}), retrying")
            time.sleep(60)
            continue
        if job.state.name in COMPLETED:
            break
        print(f"{label}: {job.state.name} ...")
        time.sleep(60)
    if job.state.name != "JOB_STATE_SUCCEEDED":
        raise SystemExit(f"{label} ended in {job.state.name} — resume with --start")
    return collect(client, job, chunk)


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--collect", metavar="JOB_NAME",
                    help="collect an existing batch job (single-batch runs only)")
    ap.add_argument("--limit", type=int, default=None,
                    help="only submit the first N prompts (test runs)")
    ap.add_argument("--stride", type=int, default=None,
                    help="take every Nth prompt (diverse test slice); use the SAME "
                         "flags when resuming with --collect")
    ap.add_argument("--start", type=int, default=0,
                    help="skip the first N prompts (resume a chunked run mid-way)")
    ap.add_argument("--chunk-size", type=int, default=None,
                    help="submit sequentially in chunks of N prompts to stay under "
                         "the batch enqueued-token quota")
    args = ap.parse_args()

    client = genai.Client()
    prompts = load_prompts()
    if args.stride:
        prompts = prompts[::args.stride]
    if args.start:
        prompts = prompts[args.start:]
    if args.limit:
        prompts = prompts[:args.limit]

    if args.collect:
        job = client.batches.get(name=args.collect)
        while job.state.name not in COMPLETED:
            print(f"state: {job.state.name} ...")
            time.sleep(60)
            try:
                job = client.batches.get(name=args.collect)
            except Exception as e:  # noqa: BLE001
                print(f"poll failed ({type(e).__name__}), retrying")
        if job.state.name != "JOB_STATE_SUCCEEDED":
            raise SystemExit(f"Batch ended in {job.state.name}")
        print(f"Done: {collect(client, job, prompts)} analyses appended to {OUT}")
        return

    size = args.chunk_size or len(prompts)
    total = 0
    for c, i in enumerate(range(0, len(prompts), size)):
        chunk = prompts[i:i + size]
        total += run_batch(client, chunk, f"deepfeline-teacher-c{c}")
        print(f"progress: {total} analyses collected, "
              f"{len(prompts) - i - len(chunk)} prompts remaining")
    print(f"\nDone: {total} analyses appended to {OUT}")


if __name__ == "__main__":
    main()
