"""research -> brief -> distilled model -> validated pack.

Talks to Ollama on localhost only. Nothing leaves the machine at inference time,
which was the point of distilling in the first place.

Ollama is given the pack JSON schema directly, so malformed output is structurally
impossible and the retry below only ever fires for empty or over-length fields.

  python serve/generate.py --niche printing3d --topic "bed adhesion"
  python serve/generate.py --niche finance --variant 2 --json
  python serve/generate.py --brief my_brief.json
"""
import argparse
import json
import pathlib
import sys
import textwrap

import requests

sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent.parent))
from common import (NICHES, PACK_SCHEMA, PACK_SYSTEM, render_brief,  # noqa: E402
                    validate_pack)
sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent))
from research import build_brief, load, retrieve  # noqa: E402

OLLAMA = "http://localhost:11434/api/chat"


def ask(model, messages, temperature, timeout=300):
    r = requests.post(OLLAMA, timeout=timeout, json={
        "model": model,
        "messages": messages,
        "format": PACK_SCHEMA,
        "stream": False,
        "options": {"temperature": temperature, "num_ctx": 4096},
    })
    if r.status_code == 404:
        sys.exit(f"Ollama has no model named {model!r}. "
                 f"Build it with: ollama create {model} -f serve/Modelfile")
    r.raise_for_status()
    return r.json()["message"]["content"]


def generate(model, brief, niche_label, temperature):
    user = render_brief(brief, niche_label)
    messages = [{"role": "system", "content": PACK_SYSTEM},
                {"role": "user", "content": user}]

    raw = ask(model, messages, temperature)
    pack = json.loads(raw)
    problems = validate_pack(pack)
    if not problems:
        return pack, []

    # One repair pass. Beyond that, hand the problems back rather than looping —
    # a model that can't satisfy the constraints twice won't on the third try.
    messages += [
        {"role": "assistant", "content": raw},
        {"role": "user", "content": "Fix these problems and return the whole object again:\n"
                                    + "\n".join(f"- {p}" for p in problems)},
    ]
    raw = ask(model, messages, temperature)
    pack = json.loads(raw)
    return pack, validate_pack(pack)


def show(pack, brief):
    w = lambda s, ind="  ": textwrap.fill(str(s), 92, initial_indent=ind,
                                          subsequent_indent=ind)
    print(f"\nBRIEF: {brief.get('topic','')}")
    print(f"ANGLE: {brief.get('angle','')}\n")
    print("HOOK")
    print(w(pack.get("hook", "")))
    print("\nSCRIPT")
    for line in str(pack.get("script", "")).split("\n"):
        if line.strip():
            print(w(line))
    yt, ig, pin = pack.get("youtube", {}), pack.get("instagram", {}), pack.get("pinterest", {})
    print(f"\nYOUTUBE  ({len(yt.get('title',''))}/100 chars)")
    print(w(yt.get("title", "")))
    print(w(yt.get("description", "")))
    print(w("tags: " + ", ".join(yt.get("tags", []))))
    print(f"\nINSTAGRAM  ({len(ig.get('caption',''))}/2200 chars)")
    for line in str(ig.get("caption", "")).split("\n"):
        print(w(line) if line.strip() else "")
    print(w(" ".join(ig.get("hashtags", []))))
    print(f"\nPINTEREST  ({len(pin.get('title',''))}/100 chars)")
    print(w(pin.get("title", "")))
    print(w(pin.get("description", "")))
    print("\nLONG CAPTION")
    for line in str(pack.get("long_caption", "")).split("\n"):
        print(w(line) if line.strip() else "")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--niche", choices=list(NICHES))
    ap.add_argument("--topic", default="")
    ap.add_argument("--brief", help="path to a JSON brief, skipping retrieval")
    ap.add_argument("--model", default="socialmedia")
    ap.add_argument("--variant", type=int, default=0)
    ap.add_argument("--k", type=int, default=8)
    ap.add_argument("--temperature", type=float, default=0.8)
    ap.add_argument("--json", action="store_true", help="raw JSON instead of a readable pack")
    args = ap.parse_args()

    if args.brief:
        brief = json.loads(pathlib.Path(args.brief).read_text())
        niche_label = NICHES.get(args.niche, "")
    else:
        if not args.niche:
            sys.exit("Pass --niche (or --brief with your own).")
        winners, briefs = load(args.niche)
        hits = retrieve(winners, briefs, args.topic, args.k)
        brief = build_brief(hits, args.topic, args.variant)
        niche_label = NICHES[args.niche]

    pack, problems = generate(args.model, brief, niche_label, args.temperature)

    if args.json:
        print(json.dumps({"brief": brief, "pack": pack, "problems": problems},
                         indent=2, ensure_ascii=False))
    else:
        show(pack, brief)
        if problems:
            print("\nstill invalid after one repair pass:")
            for p in problems:
                print(f"  - {p}")


if __name__ == "__main__":
    main()
