from __future__ import annotations

from pathlib import Path
import argparse
import json
import sqlite3
import sys
from collections import defaultdict

ROOT = Path(__file__).resolve().parents[1]
SRC = ROOT / "src"
sys.path.insert(0, str(SRC))

from ai_archive.paths import INDEX_DIR

DOMAINS = {
    "barely-but-here": {
        "title": "Barely, But Here",
        "project": "Barely, But Here",
        "queries": [
            "Barely, But Here",
            "Barely But Here",
            "not-fun places",
            "Hard-Won Wisdom",
            "Functional Depression",
            "Substack",
            "BBH",
            "content warning",
        ],
    },
    "bert-identity-voice": {
        "title": "Bert identity / professional positioning / writing voice",
        "project": "Bert identity and writing voice",
        "queries": [
            "writing voice",
            "staccato rhythm",
            "California-casual",
            "professional positioning",
            "AI product designer",
            "AI systems architect",
            "raw honesty",
            "no self-help framing",
        ],
    },
    "ai-systems-assessment": {
        "title": "AI Systems Assessment Campaign",
        "project": "AI Systems Assessment Campaign",
        "queries": [
            "AI Systems Assessment",
            "AI Assessment",
            "7-day AI Assessment",
            "$599",
            "Tools ≠ System",
            "system problem",
            "10–20 hours",
            "landing page",
        ],
    },
}


def quote_fts_query(query: str) -> str:
    return '"' + query.replace('"', '""') + '"'


def parse_md(raw: str | None) -> dict:
    try:
        return json.loads(raw or "{}")
    except Exception:
        return {"raw_metadata_json": raw}


def search(conn: sqlite3.Connection, query: str, limit: int) -> list[dict]:
    sql = """
    SELECT c.chunk_id, c.conversation_id, c.source, c.title, c.message_start, c.message_end,
           c.text, c.metadata_json,
           snippet(chunks_fts, 4, '[', ']', '…', 24) AS snippet
    FROM chunks_fts f
    JOIN chunks c ON c.rowid = f.rowid
    WHERE chunks_fts MATCH ?
    LIMIT ?
    """
    rows = conn.execute(sql, [quote_fts_query(query), limit]).fetchall()
    out = []
    for row in rows:
        d = dict(row)
        d["metadata"] = parse_md(d.pop("metadata_json"))
        d["matched_query"] = query
        out.append(d)
    return out


def candidate_type(domain: str, text: str, query: str) -> str:
    hay = f"{query}\n{text}".lower()
    if domain == "barely-but-here":
        if any(x in hay for x in ["draft", "# ", "functional depression", "hypervigilance", "substack"]):
            return "source_excerpt"
        if any(x in hay for x in ["voice", "staccato", "not-fun", "tagline", "content warning"]):
            return "writing_voice"
        return "project_context"
    if domain == "bert-identity-voice":
        if any(x in hay for x in ["voice", "staccato", "tone", "style", "raw honesty", "self-help"]):
            return "writing_voice"
        return "identity"
    if domain == "ai-systems-assessment":
        if any(x in hay for x in ["$599", "assessment", "offer", "landing page"]):
            return "offer"
        return "business_strategy"
    return "source_excerpt"


def sensitivity(text: str) -> str:
    hay = text.lower()
    if any(x in hay for x in ["therapy", "depression", "mental health", "trauma", "homelessness", "divorce"]):
        return "sensitive"
    return "private"


def make_review_packet(domain_key: str, per_query_limit: int) -> tuple[Path, int]:
    domain = DOMAINS[domain_key]
    conn = sqlite3.connect(INDEX_DIR / "archive.db")
    conn.row_factory = sqlite3.Row
    seen = set()
    buckets = defaultdict(list)
    try:
        for q in domain["queries"]:
            for row in search(conn, q, per_query_limit):
                if row["chunk_id"] in seen:
                    continue
                seen.add(row["chunk_id"])
                text = row["text"]
                md = row["metadata"]
                candidate = {
                    "candidate_id": f"{domain_key}:{row['chunk_id']}",
                    "review_status": "candidate",
                    "recommended_memory_type": candidate_type(domain_key, text, row["matched_query"]),
                    "recommended_project": domain["project"],
                    "recommended_sensitivity": sensitivity(text),
                    "matched_query": row["matched_query"],
                    "source": {
                        "source": row["source"],
                        "title": row["title"],
                        "conversation_id": row["conversation_id"],
                        "chunk_id": row["chunk_id"],
                        "raw_path": md.get("raw_path"),
                        "message_range": f"{row['message_start']}-{row['message_end']}",
                    },
                    "snippet": (row.get("snippet") or "").replace("\n", " "),
                    "candidate_content_preview": text[:1200],
                    "review_notes": "",
                }
                buckets[candidate["recommended_memory_type"]].append(candidate)
    finally:
        conn.close()

    out = ROOT / "review" / f"{domain_key}-memory-candidates.md"
    out.parent.mkdir(parents=True, exist_ok=True)
    lines = [
        f"# Memory Candidate Review Packet: {domain['title']}",
        "",
        "These are candidates only. Nothing in this packet is approved or installed until Bert reviews it.",
        "",
        "## Review legend",
        "",
        "- `APPROVE` — install or export as approved memory",
        "- `EDIT` — revise before approval",
        "- `REJECT` — keep searchable in archive only",
        "- `SOURCE_ONLY` — important source document/excerpt, but not durable memory",
        "",
        f"Total candidates: {sum(len(v) for v in buckets.values())}",
        "",
    ]
    for memory_type, items in sorted(buckets.items()):
        lines += [f"## {memory_type}", ""]
        for idx, c in enumerate(items, 1):
            src = c["source"]
            lines += [
                f"### {idx}. `{c['candidate_id']}`",
                "",
                "**Review decision:** `PENDING`  ",
                f"**Matched query:** `{c['matched_query']}`  ",
                f"**Recommended type:** `{c['recommended_memory_type']}`  ",
                f"**Recommended sensitivity:** `{c['recommended_sensitivity']}`  ",
                f"**Source:** `{src['source']}` / `{src['title']}`  ",
                f"**Raw path:** `{src['raw_path']}`  ",
                f"**Conversation:** `{src['conversation_id']}`  ",
                f"**Chunk:** `{src['chunk_id']}`  ",
                f"**Message range:** `{src['message_range']}`  ",
                "",
                "**Snippet**",
                "",
                f"> {c['snippet'][:900]}",
                "",
                "**Candidate content preview**",
                "",
                "```text",
                c["candidate_content_preview"].replace("```", "'''"),
                "```",
                "",
            ]
    out.write_text("\n".join(lines), encoding="utf-8")
    return out, sum(len(v) for v in buckets.values())


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--domain", choices=list(DOMAINS) + ["all"], default="all")
    parser.add_argument("--per-query-limit", type=int, default=8)
    args = parser.parse_args()
    domains = list(DOMAINS) if args.domain == "all" else [args.domain]
    for d in domains:
        path, count = make_review_packet(d, args.per_query_limit)
        print(f"wrote {count} candidates -> {path}")
    return 0


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