"""The research half of the pipeline — deliberately not part of the model.

Teaching a 4B model to browse is expensive and fragile. Retrieval over a corpus of
posts that already demonstrably performed is cheap, deterministic, and better
grounded. So this module is plain Python, runs entirely offline, and the model is
left to do the one thing it was distilled for: writing.

The knowledge base is free. During backtranslation the teacher already extracted
3-5 concrete factual points for every winning post, so data/briefs.jsonl is a
niche-specific research index that cost nothing extra to build.

  python serve/research.py --niche printing3d --topic "bed adhesion"
  python serve/research.py --niche finance --list      # what's working right now
"""
import argparse
import json
import sys
from collections import Counter

sys.path.insert(0, str(__import__("pathlib").Path(__file__).resolve().parent.parent))
from common import DATA, NICHES, read_jsonl, words  # noqa: E402


def load(niche):
    briefs = {b["id"]: b for b in read_jsonl(DATA / "briefs.jsonl")}
    winners = [r for r in read_jsonl(DATA / "curated.jsonl")
               if r["niche"] == niche and r["id"] in briefs]
    if not winners:
        sys.exit(f"No corpus for niche {niche!r}. Run the collectors, build/score.py "
                 f"and build/backtranslate.py first.")
    return winners, briefs


def retrieve(winners, briefs, topic, k):
    """Rank winners by topical overlap, then by how far they beat their creator's
    baseline. With no topic, engagement score alone decides."""
    q = words(topic)
    ranked = []
    for w in winners:
        b = briefs[w["id"]]
        hay = words(" ".join([
            b["brief"].get("topic", ""), b["brief"].get("angle", ""),
            " ".join(b["brief"].get("research_context") or []),
            (w.get("text") or {}).get("title", "") or (w.get("text") or {}).get("caption", ""),
        ]))
        overlap = len(q & hay) / len(q) if q else 0.0
        ranked.append((overlap, w.get("score", 0.0), w, b))
    ranked.sort(key=lambda t: (round(t[0], 3), t[1]), reverse=True)
    if q and ranked and ranked[0][0] == 0.0:
        print(f"  (nothing in the corpus matches {topic!r} — falling back to "
              f"top performers for the niche)", file=sys.stderr)
    return ranked[:k]


def build_brief(hits, topic, variant):
    """Assemble a brief from retrieved evidence.

    The angle is borrowed from one of the retrieved winners rather than invented,
    and --variant rotates which one, so repeated runs on the same topic explore
    different framings instead of returning the same post forever.
    """
    briefs = [b["brief"] for _, _, _, b in hits]
    pick = briefs[variant % len(briefs)]

    context, seen = [], set()
    for b in briefs:
        for c in (b.get("research_context") or []):
            key = " ".join(sorted(words(c)))
            if key and key not in seen:
                seen.add(key)
                context.append(c)
    audience = Counter(b.get("audience", "") for b in briefs if b.get("audience"))
    fmt = Counter(b.get("format_note", "") for b in briefs if b.get("format_note"))

    return {
        "topic": topic or pick.get("topic", ""),
        "audience": audience.most_common(1)[0][0] if audience else "",
        "angle": pick.get("angle", ""),
        "research_context": context[:6],
        "format_note": fmt.most_common(1)[0][0] if fmt else "",
    }


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--niche", choices=list(NICHES), required=True)
    ap.add_argument("--topic", default="")
    ap.add_argument("--k", type=int, default=8, help="winners to retrieve")
    ap.add_argument("--variant", type=int, default=0,
                    help="rotate which retrieved angle is used")
    ap.add_argument("--list", action="store_true",
                    help="show the retrieved winners instead of a brief")
    args = ap.parse_args()

    winners, briefs = load(args.niche)
    hits = retrieve(winners, briefs, args.topic, args.k)

    if args.list:
        print(f"top {len(hits)} for {args.niche}"
              + (f" matching {args.topic!r}" if args.topic else "") + ":\n")
        for overlap, score, w, b in hits:
            t = w.get("text") or {}
            title = t.get("title") or (t.get("caption") or "").split("\n")[0]
            print(f"  [{w['platform']:<9} score {score:+.2f}  match {overlap:.0%}] {title[:80]}")
            print(f"      hook: {b['synth'].get('hook','')[:90]}")
            print(f"      why : {b.get('why_it_works','')[:90]}")
            print(f"      {w.get('url','')}\n")
        return

    print(json.dumps(build_brief(hits, args.topic, args.variant),
                     indent=2, ensure_ascii=False))


if __name__ == "__main__":
    main()
