#!/usr/bin/env python3
"""Retrieval smoke test for the per-agent KB RAG pipeline.

Usage:
  source .venv/bin/activate
  python scripts/test_kb_retrieval.py [path/to/file.pdf]

If no PDF is given, uses a built-in text fixture with known facts.
"""

from __future__ import annotations

import sys
import tempfile
import uuid
from pathlib import Path

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

from app.config import get_settings

get_settings.cache_clear()

from app.saas import db as saas_db
from app.saas.kb import (
    agent_collection_name,
    agent_document_count,
    drop_agent_store,
    get_agent_vectorstore,
    ingest_kb_file,
    org_storage_dir,
    retrieve_agent_scored,
)


FIXTURE_TEXT = """\
Acme Finance Policy Handbook

Refund Policy
Customers may request a full refund within 14 days of purchase.
Late fees are 2 percent per month on overdue invoices.
Contact support at support@acme.test for billing questions.

Payment Terms
Standard net payment terms are Net 30.
Wire transfers must include the invoice number in the reference field.
"""

KNOWN_QUERIES = [
    ("How many days is the refund window?", ["14 days", "refund"]),
    ("What is the late fee?", ["2 percent", "late"]),
    ("What are the payment terms?", ["Net 30", "payment"]),
    ("What email do I use for billing support?", ["support@acme.test"]),
]


def _ensure_fixture_pdf_or_txt(path: Path | None) -> tuple[Path, list[tuple[str, list[str]]]]:
    if path and path.exists():
        # For real PDFs, only check that retrieval returns non-empty hits.
        return path, [
            ("What is this document about?", []),
            ("Summarize the totals or main topic", []),
            ("List key amounts or policies mentioned", []),
        ]

    tmp = Path(tempfile.mkdtemp()) / "acme_policy.txt"
    tmp.write_text(FIXTURE_TEXT, encoding="utf-8")
    return tmp, KNOWN_QUERIES


def main() -> int:
    settings = get_settings()
    print("Embedding model:", settings.embedding_model)
    print("KB top_k:", settings.kb_retrieval_top_k)
    print("KB threshold:", settings.kb_similarity_threshold)
    print("Chunk size/overlap:", settings.kb_chunk_size, settings.kb_chunk_overlap)
    print("RAG_DEBUG:", settings.rag_debug)
    print()

    arg = Path(sys.argv[1]) if len(sys.argv) > 1 else None
    src, queries = _ensure_fixture_pdf_or_txt(arg)

    saas_db.init_saas_tables()
    org_id = f"test-org-{uuid.uuid4()}"
    agent_id = f"test-agent-{uuid.uuid4()}"
    # Minimal org/agent rows are not strictly required for Chroma, but doc table is.
    doc = saas_db.create_kb_document(
        org_id=org_id,
        filename=src.name,
        content_type="text/plain" if src.suffix == ".txt" else "application/pdf",
        uploaded_by=None,
        agent_id=agent_id,
    )
    dest = org_storage_dir(org_id, agent_id) / f"{doc['id']}_{src.name}"
    dest.write_bytes(src.read_bytes())

    print("Indexing:", dest)
    try:
        n = ingest_kb_file(org_id, doc["id"], dest, src.name, agent_id=agent_id)
    except Exception as exc:
        print("FAIL ingest:", exc)
        return 1

    store = get_agent_vectorstore(org_id, agent_id)
    count = agent_document_count(org_id, agent_id)
    print(f"Indexed chunks={n} chroma_count={count} collection={agent_collection_name(agent_id)}")
    if count <= 0:
        print("FAIL: Chroma collection empty after ingest")
        return 1

    sample = store._collection.get(limit=1, include=["documents", "metadatas"])
    if sample.get("documents"):
        print("Example metadata:", sample["metadatas"][0])
        print("Example doc:", (sample["documents"][0] or "")[:200].replace("\n", " "))
    print()

    failures = 0
    for question, needles in queries:
        scored, reason = retrieve_agent_scored(org_id, agent_id, question)
        print("=" * 60)
        print("Query:", question)
        if reason:
            print("Empty reason:", reason)
        if not scored:
            print("FAIL: no hits")
            failures += 1
            continue
        blob = "\n".join(d.page_content for d, _ in scored).lower()
        for i, (doc, score) in enumerate(scored, start=1):
            meta = doc.metadata or {}
            print(
                f"  [{i}] score={score} page={meta.get('page')} "
                f"id={meta.get('record_id')} file={meta.get('source_dataset')}"
            )
            print("      ", doc.page_content[:180].replace("\n", " "))
        if needles:
            missing = [n for n in needles if n.lower() not in blob]
            if missing:
                print("FAIL: expected terms not in retrieved context:", missing)
                failures += 1
            else:
                print("PASS: expected terms found in retrieved context")
        else:
            print("PASS: retrieved", len(scored), "chunk(s)")

    drop_agent_store(org_id, agent_id)
    print()
    if failures:
        print(f"RESULT: {failures} failure(s)")
        return 1
    print("RESULT: all retrieval checks passed")
    return 0


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