"""Per-agent knowledge base: extract → chunk → embed → Chroma → retrieve."""

from __future__ import annotations

import hashlib
import re
from pathlib import Path
from typing import Any

from langchain_community.vectorstores import Chroma
from langchain_core.documents import Document

from app.config import BASE_DIR, get_settings
from app.rag.chain import get_rag_service
from app.saas import db as saas_db
from app.saas import rag_debug as dbg

KB_UPLOAD_DIR = BASE_DIR / "data" / "kb_uploads"
AVATAR_DIR = BASE_DIR / "data" / "agent_avatars"
TENANT_CHROMA_ROOT = BASE_DIR / "data" / "chroma" / "tenants"

MIN_CHUNK_CHARS = 40
MIN_DOC_CHARS = 20  # allow short notes as a single chunk
MAX_CHUNK_CHARS = 6000  # hard cap; settings drive normal size
EMPTY_EXTRACT_WARN_CHARS = 80
EMBED_BATCH = 16

_agent_stores: dict[str, Chroma] = {}
_embed_dim_cache: int | None = None


def agent_collection_name(agent_id: str) -> str:
    return f"a_{agent_id.replace('-', '')[:24]}"


def agent_persist_dir(org_id: str, agent_id: str) -> Path:
    path = TENANT_CHROMA_ROOT / org_id / agent_id
    path.mkdir(parents=True, exist_ok=True)
    return path


def get_agent_vectorstore(org_id: str, agent_id: str) -> Chroma:
    key = f"{org_id}:{agent_id}"
    if key in _agent_stores:
        return _agent_stores[key]
    service = get_rag_service()
    name = agent_collection_name(agent_id)
    persist = str(agent_persist_dir(org_id, agent_id))
    store = Chroma(
        collection_name=name,
        embedding_function=service.embeddings,
        persist_directory=persist,
    )
    _agent_stores[key] = store
    dbg.log_info(
        "Opened agent Chroma collection=%s persist=%s embed_model=%s",
        name,
        persist,
        get_settings().embedding_model,
    )
    return store


def agent_document_count(org_id: str, agent_id: str) -> int:
    try:
        return get_agent_vectorstore(org_id, agent_id)._collection.count()
    except Exception as exc:
        dbg.log_warn("agent_document_count failed: %s", exc)
        return 0


def get_tenant_vectorstore(org_id: str) -> Chroma:
    path = TENANT_CHROMA_ROOT / org_id
    path.mkdir(parents=True, exist_ok=True)
    key = f"org:{org_id}"
    if key in _agent_stores:
        return _agent_stores[key]
    service = get_rag_service()
    store = Chroma(
        collection_name=f"org_{org_id.replace('-', '')[:24]}",
        embedding_function=service.embeddings,
        persist_directory=str(path),
    )
    _agent_stores[key] = store
    return store


def tenant_document_count(org_id: str) -> int:
    try:
        return get_tenant_vectorstore(org_id)._collection.count()
    except Exception:
        return 0


def normalize_whitespace(text: str) -> str:
    text = (text or "").replace("\u00a0", " ")
    text = re.sub(r"[ \t]+", " ", text)
    text = re.sub(r"\n{3,}", "\n\n", text)
    return text.strip()


def _is_heading(line: str) -> bool:
    s = line.strip()
    if not s or len(s) > 120:
        return False
    if s.startswith("#"):
        return True
    if re.match(r"^\d+(\.\d+)*\s+\S", s):
        return True
    letters = [c for c in s if c.isalpha()]
    if len(letters) >= 4 and sum(1 for c in letters if c.isupper()) / len(letters) > 0.75:
        return True
    return False


def _is_list_item(line: str) -> bool:
    return bool(re.match(r"^([-*•]|\d+[.)])\s+\S", line.strip()))


def _is_table_row(line: str) -> bool:
    s = line.strip()
    return s.count("|") >= 2 or s.count("\t") >= 2


def _split_structural_units(text: str) -> list[str]:
    """Split text into headings, paragraphs, list blocks, and table blocks."""
    units: list[str] = []
    buf: list[str] = []
    mode: str | None = None  # 'para' | 'list' | 'table'

    def flush() -> None:
        nonlocal mode
        if buf:
            units.append("\n".join(buf).strip())
            buf.clear()
        mode = None

    for raw in (text or "").splitlines():
        line = raw.rstrip()
        stripped = line.strip()
        if not stripped:
            flush()
            continue
        if _is_heading(stripped):
            flush()
            units.append(stripped)
            continue
        if _is_table_row(stripped):
            if mode not in (None, "table"):
                flush()
            mode = "table"
            buf.append(stripped)
            continue
        if _is_list_item(stripped):
            if mode not in (None, "list"):
                flush()
            mode = "list"
            buf.append(stripped)
            continue
        if mode in ("list", "table"):
            flush()
        mode = "para"
        buf.append(stripped)
    flush()
    return [u for u in units if u]


def _hard_wrap(text: str, chunk_size: int, overlap: int) -> list[str]:
    pieces: list[str] = []
    start = 0
    n = len(text)
    while start < n:
        end = min(start + chunk_size, n)
        if end < n:
            window = text[start:end]
            br = max(window.rfind("\n"), window.rfind(". "), window.rfind(" "))
            if br > chunk_size // 3:
                end = start + br + 1
        piece = text[start:end].strip()
        if piece:
            pieces.append(piece)
        if end >= n:
            break
        start = max(0, end - overlap)
    return pieces


def chunk_text(
    text: str,
    chunk_size: int | None = None,
    overlap: int | None = None,
) -> list[str]:
    """Structure-aware chunking with overlap between consecutive chunks."""
    settings = get_settings()
    chunk_size = chunk_size or settings.kb_chunk_size
    overlap = overlap or settings.kb_chunk_overlap
    text = normalize_whitespace(text)
    if not text:
        return []
    if len(text) < MIN_DOC_CHARS:
        return []
    if len(text) <= chunk_size:
        return [text[:MAX_CHUNK_CHARS]]

    units = _split_structural_units(text)
    if not units:
        units = [text]

    packed: list[str] = []
    buf = ""
    for unit in units:
        unit = unit.strip()
        if not unit:
            continue
        if len(unit) > chunk_size:
            if buf:
                packed.append(buf)
                buf = ""
            packed.extend(_hard_wrap(unit, chunk_size, overlap))
            continue
        if not buf:
            buf = unit
        elif len(buf) + 2 + len(unit) <= chunk_size:
            buf = f"{buf}\n\n{unit}"
        else:
            packed.append(buf)
            buf = unit
    if buf:
        packed.append(buf)

    # Merge undersized trailing fragments into previous chunk when possible
    merged: list[str] = []
    for piece in packed:
        piece = piece.strip()
        if not piece:
            continue
        if (
            merged
            and len(piece) < MIN_CHUNK_CHARS
            and len(merged[-1]) + 2 + len(piece) <= chunk_size
        ):
            merged[-1] = f"{merged[-1]}\n\n{piece}"
        else:
            merged.append(piece[:MAX_CHUNK_CHARS])

    if not merged and len(text) >= MIN_DOC_CHARS:
        merged = [text[:MAX_CHUNK_CHARS]]

    # Apply overlap: prefix each chunk (except first) with a trailing window of the previous
    if overlap > 0 and len(merged) > 1:
        overlapped: list[str] = [merged[0]]
        for i in range(1, len(merged)):
            prev = merged[i - 1]
            tail = prev[-overlap:]
            # Prefer word/line boundary
            sp = max(tail.find("\n"), tail.find(" "))
            if 0 < sp < len(tail) // 2:
                tail = tail[sp + 1 :]
            next_chunk = merged[i]
            if not next_chunk.startswith(tail.strip()[:40]):
                combo = f"{tail.strip()}\n\n{next_chunk}".strip()
                overlapped.append(combo[:MAX_CHUNK_CHARS])
            else:
                overlapped.append(next_chunk[:MAX_CHUNK_CHARS])
        merged = overlapped

    out = [c.strip() for c in merged if c and len(c.strip()) >= min(MIN_CHUNK_CHARS, MIN_DOC_CHARS)]
    return out


def _meta(
    org_id: str,
    agent_id: str,
    doc_id: str,
    filename: str,
    page: int,
    rid: str,
    *,
    chunk_index: int = 0,
    chunk_kind: str | None = None,
    content: str = "",
    parent_chunk_id: str | None = None,
    section: str | None = None,
) -> dict:
    kind = chunk_kind or infer_chunk_kind(content)
    meta = {
        "record_id": f"kb_{rid}",
        "chunk_id": f"kb_{rid}",
        "org_id": org_id,
        "agent_id": agent_id,
        "doc_id": doc_id,
        "data_category": "tenant_kb",
        "source_dataset": filename,
        "source_file": filename,
        "page": int(page),
        "chunk_index": int(chunk_index),
        "chunk_kind": kind,
    }
    # Hierarchical retrieval: child chunks point at a page/section parent id.
    if parent_chunk_id:
        meta["parent_chunk_id"] = parent_chunk_id
    if section:
        meta["section"] = section
    return meta


def infer_chunk_kind(text: str) -> str:
    """Classify chunk content for metadata-filtered retrieval."""
    t = (text or "").lower()
    if re.search(
        r"statement\s+period|from\s+date|to\s+date|period\s*:|date\s*range",
        t,
    ):
        return "statement_period"
    if re.search(
        r"account\s+holder|customer\s+name|a/?c\s*(?:no|number)|ifsc|micr|"
        r"branch\s*(?:code|name)|cust(?:omer)?\s*id|opening\s+balance|closing\s+balance",
        t,
    ):
        return "header"
    if re.search(
        r"total\s+money|payments?\s+made|money\s+paid|money\s+received|"
        r"total\s+(?:spent|debit|credit)|net\s+cash",
        t,
    ):
        return "totals"
    if re.search(r"\bsummary\b|\boverview\b", t) and len(t) < 1200:
        return "summary"
    if re.search(r"merchant|upi\s+id|paid\s+to|vpa", t):
        return "merchant"
    if re.search(
        r"₹|rs\.?\s*\d|\d{1,2}\s+(?:jan|feb|mar|apr|may|jun|jul|aug|sep|oct|nov|dec)|"
        r"debit|credit|txn|transaction",
        t,
    ):
        return "transaction"
    return "summary"


def _probe_embedding_dim() -> int:
    global _embed_dim_cache
    if _embed_dim_cache is not None:
        return _embed_dim_cache
    service = get_rag_service()
    vec = service.embeddings.embed_query("probe")
    if not vec:
        raise RuntimeError("Embedding generation returned empty vector")
    _embed_dim_cache = len(vec)
    dbg.log_info(
        "Embedding model=%s dimension=%s",
        get_settings().embedding_model,
        _embed_dim_cache,
    )
    return _embed_dim_cache


def _set_progress(
    doc_id: str,
    *,
    stage: str,
    pct: int,
    status: str = "processing",
    error: str | None = None,
    page_count: int | None = None,
    pages_processed: int | None = None,
    chunks_created: int | None = None,
    embeddings_done: int | None = None,
    vectors_stored: int | None = None,
    chunk_count: int | None = None,
) -> None:
    saas_db.update_kb_document(
        doc_id,
        status=status,
        progress_stage=stage,
        progress_pct=max(0, min(100, int(pct))),
        error=error,
        page_count=page_count,
        pages_processed=pages_processed,
        chunks_created=chunks_created,
        embeddings_done=embeddings_done,
        vectors_stored=vectors_stored,
        chunk_count=chunk_count,
    )


def extract_pdf_pages(path: Path) -> tuple[list[tuple[int, str]], dict]:
    """Return [(page_num, text), ...] for pages with text, plus extraction stats.

    Every page is examined; empty/scanned pages are counted but omitted from text pairs.
    """
    import pdfplumber

    pages: list[tuple[int, str]] = []
    empty_pages = 0
    with pdfplumber.open(str(path)) as pdf:
        total_pages = len(pdf.pages)
        for page_num, page in enumerate(pdf.pages, start=1):
            page_text = normalize_whitespace(page.extract_text() or "")
            # Preserve table structure when available
            try:
                tables = page.extract_tables() or []
            except Exception:
                tables = []
            table_blocks: list[str] = []
            for table in tables:
                rows: list[str] = []
                for row in table or []:
                    cells = [str(c).strip() if c is not None else "" for c in row]
                    if any(cells):
                        rows.append(" | ".join(cells))
                if rows:
                    table_blocks.append("[Table]\n" + "\n".join(rows))
            if table_blocks:
                extra = "\n\n".join(table_blocks)
                page_text = normalize_whitespace(f"{page_text}\n\n{extra}" if page_text else extra)
            if not page_text:
                empty_pages += 1
                continue
            pages.append((page_num, page_text))

    full = "\n\n".join(t for _, t in pages)
    stats = {
        "total_pages": total_pages,
        "pages_with_text": len(pages),
        "pages_processed": total_pages,  # all pages examined
        "empty_pages": empty_pages,
        "total_chars": len(full),
        "preview_500": full[:500],
        "likely_scanned": total_pages > 0 and len(pages) == 0,
        "very_short": 0 < len(full) < EMPTY_EXTRACT_WARN_CHARS,
        "pages_skipped": 0,  # we examine every page
    }
    if stats["pages_processed"] != stats["total_pages"]:
        raise RuntimeError(
            f"Page coverage mismatch: processed={stats['pages_processed']} total={stats['total_pages']}"
        )
    return pages, stats


def _docs_from_pdf(
    path: Path, org_id: str, agent_id: str, doc_id: str, filename: str
) -> tuple[list[Document], dict]:
    pages, stats = extract_pdf_pages(path)
    dbg.log_info(
        "PDF extract file=%s pages=%s/%s chars=%s scanned=%s",
        filename,
        stats["pages_with_text"],
        stats["total_pages"],
        stats["total_chars"],
        stats["likely_scanned"],
    )
    dbg.log_info("PDF extract preview (500 chars): %s", stats["preview_500"])
    if stats["likely_scanned"]:
        dbg.log_warn(
            "PDF appears scanned/image-only (no extractable text): %s — OCR not applied in KB ingest",
            filename,
        )
        raise ValueError(
            "No extractable text in PDF (likely scanned/image). "
            "Export a text PDF or OCR it before upload."
        )
    if stats["very_short"]:
        dbg.log_warn(
            "PDF extraction very short (%s chars) for %s",
            stats["total_chars"],
            filename,
        )
    if not pages:
        raise ValueError("PDF text extraction returned empty content")

    docs: list[Document] = []
    seen_hashes: set[str] = set()
    global_idx = 0
    for page_num, page_text in pages:
        # Parent id for hierarchical retrieval (page → child chunks).
        parent_rid = hashlib.md5(
            f"{agent_id}:{doc_id}:parent:p{page_num}".encode()
        ).hexdigest()[:16]
        parent_chunk_id = f"kb_{parent_rid}"
        headed = f"[Page {page_num} | {filename}]\n{page_text}"
        for idx, chunk in enumerate(chunk_text(headed)):
            digest = hashlib.md5(chunk.encode("utf-8", errors="ignore")).hexdigest()
            if digest in seen_hashes:
                dbg.log_warn("Skipping duplicate chunk page=%s idx=%s", page_num, idx)
                continue
            seen_hashes.add(digest)
            rid = hashlib.md5(
                f"{agent_id}:{doc_id}:{page_num}:{idx}:{digest[:8]}".encode()
            ).hexdigest()[:16]
            docs.append(
                Document(
                    page_content=chunk,
                    metadata=_meta(
                        org_id,
                        agent_id,
                        doc_id,
                        filename,
                        page_num,
                        rid,
                        chunk_index=global_idx,
                        content=chunk,
                        parent_chunk_id=parent_chunk_id,
                        section=f"page_{page_num}",
                    ),
                )
            )
            global_idx += 1
    stats = {**stats, "chunks_created": len(docs), "unique_chunk_hashes": len(seen_hashes)}
    return docs, stats


def _docs_from_text(
    path: Path, org_id: str, agent_id: str, doc_id: str, filename: str
) -> tuple[list[Document], dict]:
    text = normalize_whitespace(path.read_text(encoding="utf-8", errors="ignore"))
    dbg.log_info("Text extract file=%s chars=%s preview=%s", filename, len(text), text[:500])
    if len(text) < MIN_DOC_CHARS:
        raise ValueError(
            f"No usable text extracted (file too short, empty, or image-only PDF). "
            f"Need at least {MIN_DOC_CHARS} characters."
        )
    docs: list[Document] = []
    seen: set[str] = set()
    parent_rid = hashlib.md5(f"{agent_id}:{doc_id}:parent:file".encode()).hexdigest()[:16]
    parent_chunk_id = f"kb_{parent_rid}"
    for idx, chunk in enumerate(chunk_text(f"[File {filename}]\n{text}")):
        digest = hashlib.md5(chunk.encode()).hexdigest()
        if digest in seen:
            continue
        seen.add(digest)
        rid = hashlib.md5(f"{agent_id}:{doc_id}:t:{idx}:{digest[:8]}".encode()).hexdigest()[:16]
        docs.append(
            Document(
                page_content=chunk,
                metadata=_meta(
                    org_id,
                    agent_id,
                    doc_id,
                    filename,
                    1,
                    rid,
                    chunk_index=idx,
                    content=chunk,
                    parent_chunk_id=parent_chunk_id,
                    section="file",
                ),
            )
        )
    stats = {
        "total_pages": 1,
        "pages_with_text": 1 if text else 0,
        "pages_processed": 1,
        "empty_pages": 0,
        "total_chars": len(text),
        "chunks_created": len(docs),
    }
    return docs, stats


def _docs_from_csv(
    path: Path, org_id: str, agent_id: str, doc_id: str, filename: str
) -> tuple[list[Document], dict]:
    import pandas as pd

    df = pd.read_csv(path)
    dbg.log_info("CSV extract file=%s rows=%s cols=%s", filename, len(df), list(df.columns))
    docs: list[Document] = []
    seen: set[str] = set()
    for idx, row in df.iterrows():
        parts = [f"{col}: {row[col]}" for col in df.columns if str(row[col]) not in {"nan", ""}]
        content = normalize_whitespace(" | ".join(parts))
        if len(content) < MIN_CHUNK_CHARS:
            continue
        digest = hashlib.md5(content.encode()).hexdigest()
        if digest in seen:
            continue
        seen.add(digest)
        rid = hashlib.md5(f"{agent_id}:{doc_id}:csv:{idx}:{digest[:8]}".encode()).hexdigest()[:16]
        page = int(idx) + 1 if isinstance(idx, int) else 1
        docs.append(
            Document(
                page_content=content[:MAX_CHUNK_CHARS],
                metadata=_meta(
                    org_id,
                    agent_id,
                    doc_id,
                    filename,
                    page,
                    rid,
                    chunk_index=int(idx) if isinstance(idx, int) else 0,
                    content=content,
                ),
            )
        )
    stats = {
        "total_pages": max(1, len(df)),
        "pages_with_text": len(docs),
        "pages_processed": len(df),
        "empty_pages": max(0, len(df) - len(docs)),
        "total_chars": sum(len(d.page_content) for d in docs),
        "chunks_created": len(docs),
    }
    return docs, stats


def build_documents(
    path: Path, org_id: str, agent_id: str, doc_id: str, filename: str
) -> tuple[list[Document], dict]:
    suffix = path.suffix.lower()
    if suffix == ".pdf":
        docs, extract_stats = _docs_from_pdf(path, org_id, agent_id, doc_id, filename)
    elif suffix in {".txt", ".md", ".markdown"}:
        docs, extract_stats = _docs_from_text(path, org_id, agent_id, doc_id, filename)
    elif suffix == ".csv":
        docs, extract_stats = _docs_from_csv(path, org_id, agent_id, doc_id, filename)
    else:
        raise ValueError(f"Unsupported file type: {suffix}. Use PDF, TXT, MD, or CSV.")

    stats = dbg.summarize_chunks(d.page_content for d in docs)
    dbg.log_info(
        "Chunking complete file=%s count=%s avg=%s min=%s max=%s pages=%s/%s",
        filename,
        stats["count"],
        stats["avg"],
        stats["min"],
        stats["max"],
        extract_stats.get("pages_with_text"),
        extract_stats.get("total_pages"),
    )
    if stats["count"] == 0:
        raise ValueError("Chunking produced zero chunks")
    if stats["min"] < MIN_CHUNK_CHARS:
        dbg.log_warn("Some chunks are very small (min=%s)", stats["min"])
    if stats["max"] > MAX_CHUNK_CHARS:
        dbg.log_warn("Some chunks exceed hard cap (max=%s)", stats["max"])

    # Validate no duplicate IDs / missing indices
    ids = [str((d.metadata or {}).get("record_id")) for d in docs]
    if len(ids) != len(set(ids)):
        raise RuntimeError("Duplicate chunk IDs generated during chunking")
    if len(docs) != extract_stats.get("chunks_created", len(docs)):
        raise RuntimeError("Chunk count mismatch after build")

    meta = {**extract_stats, **stats, "chunks_created": len(docs)}
    return docs, meta


def _count_vectors_for_doc(store: Chroma, doc_id: str) -> int:
    try:
        got = store._collection.get(where={"doc_id": doc_id}, include=[])
        return len(got.get("ids") or [])
    except Exception as exc:
        dbg.log_warn("vector count by doc_id failed: %s", exc)
        return -1


def ingest_kb_file(
    org_id: str,
    doc_id: str,
    path: Path,
    filename: str,
    *,
    agent_id: str,
) -> int:
    _set_progress(doc_id, stage="Extracting text…", pct=5, status="processing", error="")
    try:
        dim = _probe_embedding_dim()
        _set_progress(doc_id, stage="Creating chunks…", pct=15)
        docs, meta = build_documents(path, org_id, agent_id, doc_id, filename)
        n = len(docs)
        _set_progress(
            doc_id,
            stage=f"Created {n} chunks",
            pct=25,
            page_count=int(meta.get("total_pages") or 0),
            pages_processed=int(meta.get("pages_processed") or 0),
            chunks_created=n,
        )

        # Remove any prior vectors for this doc (retry-safe)
        store = get_agent_vectorstore(org_id, agent_id)
        try:
            store._collection.delete(where={"doc_id": doc_id})
        except Exception:
            pass

        collection = agent_collection_name(agent_id)
        before = store._collection.count()
        ids = [str((d.metadata or {}).get("record_id")) for d in docs]
        if len(ids) != len(set(ids)):
            raise RuntimeError("Duplicate chunk IDs generated during ingest")

        embeddings_done = 0
        for i in range(0, n, EMBED_BATCH):
            batch_docs = docs[i : i + EMBED_BATCH]
            batch_ids = ids[i : i + EMBED_BATCH]
            pct = 25 + int(70 * (embeddings_done / max(1, n)))
            _set_progress(
                doc_id,
                stage=f"Generating embeddings… ({embeddings_done}/{n})",
                pct=pct,
                chunks_created=n,
                embeddings_done=embeddings_done,
            )
            try:
                store.add_documents(batch_docs, ids=batch_ids)
            except Exception as exc:
                dbg.log_error("Embedding/storage failed at batch %s: %s", i // EMBED_BATCH, exc)
                raise RuntimeError(f"Embedding or Chroma insert failed: {exc}") from exc
            embeddings_done += len(batch_docs)
            _set_progress(
                doc_id,
                stage=f"Generating embeddings… ({embeddings_done}/{n})",
                pct=25 + int(70 * (embeddings_done / max(1, n))),
                chunks_created=n,
                embeddings_done=embeddings_done,
            )

        _set_progress(
            doc_id,
            stage="Indexing…",
            pct=95,
            chunks_created=n,
            embeddings_done=embeddings_done,
        )

        after = store._collection.count()
        inserted = after - before
        stored = _count_vectors_for_doc(store, doc_id)
        dbg.log_info(
            "Chroma ingest collection=%s before=%s after=%s inserted≈%s "
            "chunks=%s embeddings=%s vectors_for_doc=%s dim=%s",
            collection,
            before,
            after,
            inserted,
            n,
            embeddings_done,
            stored,
            dim,
        )

        # Hard validation: every chunk embedded exactly once
        if embeddings_done != n:
            raise RuntimeError(
                f"Embedding count mismatch: embeddings_done={embeddings_done} chunks={n}"
            )
        if stored >= 0 and stored != n:
            raise RuntimeError(
                f"Vector store mismatch: stored={stored} expected={n} for doc_id={doc_id}"
            )
        if inserted < n and stored < 0:
            # Fallback when where-filter unsupported — require collection growth
            if after < before + n:
                raise RuntimeError(
                    f"Chroma count did not grow enough: before={before} after={after} expected+{n}"
                )

        sample = docs[0]
        dbg.log_info(
            "Ingest validated file=%s pages=%s/%s chunks=%s embeddings=%s vectors=%s example=%s",
            filename,
            meta.get("pages_with_text"),
            meta.get("total_pages"),
            n,
            embeddings_done,
            stored if stored >= 0 else inserted,
            (sample.page_content[:160]).replace("\n", " "),
        )

        # Dual-write structured FDI index (SQL SoR) — does not replace vectors.
        if get_settings().fdi_enabled and path.suffix.lower() == ".pdf":
            try:
                _set_progress(doc_id, stage="Indexing structured finance data…", pct=97)
                from app.fdi.pipeline import index_document_structured

                index_document_structured(
                    org_id=org_id,
                    agent_id=agent_id,
                    kb_doc_id=doc_id,
                    path=path,
                    filename=filename,
                    mime_type="application/pdf",
                )
            except Exception as fdi_exc:
                # Vectors remain usable; structured path degrades to RAG until re-index.
                dbg.log_warn("FDI structured index failed doc_id=%s: %s", doc_id, fdi_exc)

        _set_progress(
            doc_id,
            stage="Done",
            pct=100,
            status="ready",
            error="",
            page_count=int(meta.get("total_pages") or 0),
            pages_processed=int(meta.get("pages_processed") or 0),
            chunks_created=n,
            embeddings_done=embeddings_done,
            vectors_stored=stored if stored >= 0 else n,
            chunk_count=n,
        )
        agent = saas_db.get_agent(agent_id)
        if agent and agent.get("kb_mode") == "platform":
            saas_db.update_agent(agent_id, kb_mode="combined")
        return n
    except Exception as exc:
        dbg.log_error("ingest_kb_file failed doc_id=%s file=%s err=%s", doc_id, filename, exc)
        saas_db.update_kb_document(
            doc_id,
            status="failed",
            error=str(exc)[:500],
            progress_stage="Failed",
            progress_pct=0,
        )
        raise


def delete_kb_vectors(org_id: str, doc_id: str, agent_id: str | None = None) -> None:
    if not agent_id:
        doc = saas_db.get_kb_document(doc_id)
        agent_id = (doc or {}).get("agent_id")
    if not agent_id:
        return
    store = get_agent_vectorstore(org_id, agent_id)
    try:
        store._collection.delete(where={"doc_id": doc_id})
    except Exception:
        try:
            store._collection.delete(where={"doc_id": {"$eq": doc_id}})
        except Exception as exc:
            dbg.log_warn("delete_kb_vectors failed: %s", exc)


def drop_agent_store(org_id: str, agent_id: str) -> None:
    key = f"{org_id}:{agent_id}"
    _agent_stores.pop(key, None)
    path = agent_persist_dir(org_id, agent_id)
    import shutil

    if path.exists():
        shutil.rmtree(path, ignore_errors=True)


def retrieve_agent_scored(
    org_id: str,
    agent_id: str,
    question: str,
    k: int | None = None,
    *,
    chunk_kinds: tuple[str, ...] | list[str] | None = None,
) -> tuple[list[tuple[Document, float | None]], str | None]:
    """Return (scored docs, empty_reason). Never silently swallow failures.

    Optional chunk_kinds filters to header/summary/merchant/transaction/totals/
    statement_period (metadata when present; content heuristics for older indexes).
    """
    settings = get_settings()
    question = normalize_whitespace(question)
    if not question:
        return [], "empty_query"

    store = get_agent_vectorstore(org_id, agent_id)
    collection = agent_collection_name(agent_id)
    try:
        count = store._collection.count()
    except Exception as exc:
        return [], f"chroma_count_failed: {exc}"

    if count == 0:
        ready = saas_db.count_ready_kb_chunks(org_id, agent_id)
        if ready > 0:
            return [], (
                f"collection_empty but SQLite reports {ready} ready chunks — "
                f"re-upload or re-index (collection={collection})"
            )
        return [], f"collection_empty (collection={collection})"

    top_k = k or settings.kb_retrieval_top_k
    kind_filter = {k.lower() for k in (chunk_kinds or []) if k}
    # Over-fetch when filtering so we still fill top_k after kind match.
    fetch_k = max(top_k * (5 if kind_filter else 3), 20)
    threshold = settings.kb_similarity_threshold
    scored: list[tuple[Document, float | None]] = []

    try:
        scored = list(
            store.similarity_search_with_relevance_scores(question, k=min(fetch_k, count))
        )
        scored.sort(key=lambda pair: (pair[1] is None, -(pair[1] or 0.0)))
    except Exception as sim_exc:
        dbg.log_warn("similarity_search_with_relevance_scores failed (%s); trying MMR", sim_exc)
        try:
            mmr_docs = store.max_marginal_relevance_search(
                question, k=min(fetch_k, count), fetch_k=min(fetch_k, count)
            )
            scored = [(d, None) for d in mmr_docs]
        except Exception as mmr_exc:
            dbg.log_error("MMR failed: %s", mmr_exc)
            try:
                docs = store.similarity_search(question, k=min(top_k * 2, count))
                scored = [(d, None) for d in docs]
            except Exception as exc:
                return [], f"retrieval_failed: {exc}"

    if not scored:
        return [], f"search_returned_zero (collection={collection}, vectors={count})"

    kept = [(d, s) for d, s in scored if s is None or s >= threshold]
    if not kept:
        dbg.log_warn(
            "All %s hits below kb_similarity_threshold=%.3f — keeping unfiltered top results",
            len(scored),
            threshold,
        )
        kept = scored

    if kind_filter:
        filtered: list[tuple[Document, float | None]] = []
        for doc, score in kept:
            meta = doc.metadata or {}
            kind = str(meta.get("chunk_kind") or "").lower() or infer_chunk_kind(doc.page_content)
            if kind in kind_filter:
                filtered.append((doc, score))
        if filtered:
            kept = filtered
        else:
            dbg.log_warn(
                "chunk_kind filter %s matched 0 of %s hits — falling back to unfiltered",
                sorted(kind_filter),
                len(kept),
            )

    out: list[tuple[Document, float | None]] = []
    seen: set[str] = set()
    for doc, score in kept:
        rid = str(
            (doc.metadata or {}).get("record_id")
            or hashlib.md5(doc.page_content.encode()).hexdigest()
        )
        if rid in seen:
            continue
        seen.add(rid)
        out.append((doc, score))
        if len(out) >= top_k:
            break

    dbg.log_info(
        "retrieve_agent query=%r collection=%s vectors=%s returned=%s kinds=%s top_score=%s",
        question[:120],
        collection,
        count,
        len(out),
        sorted(kind_filter) if kind_filter else None,
        out[0][1] if out else None,
    )
    for i, (doc, score) in enumerate(out[:10], start=1):
        meta = doc.metadata or {}
        dbg.log_info(
            "  hit[%s] score=%s page=%s kind=%s file=%s preview=%s",
            i,
            score,
            meta.get("page"),
            meta.get("chunk_kind") or infer_chunk_kind(doc.page_content),
            meta.get("source_dataset"),
            (doc.page_content[:120]).replace("\n", " "),
        )
    if not out:
        return [], f"filtered_empty (collection={collection})"
    return out, None


def retrieve_agent(org_id: str, agent_id: str, question: str, k: int | None = None) -> list[Document]:
    scored, reason = retrieve_agent_scored(org_id, agent_id, question, k=k)
    docs = [d for d, _ in scored]
    scores = [s for _, s in scored]
    dbg.print_rag_debug(query=question, docs=docs, scores=scores, reason_if_empty=reason)
    return docs


def retrieve_tenant(org_id: str, question: str, k: int | None = None) -> list[Document]:
    # Legacy org-level store — same soft retrieval path pattern
    settings = get_settings()
    store = get_tenant_vectorstore(org_id)
    if store._collection.count() == 0:
        return []
    top_k = k or settings.kb_retrieval_top_k
    try:
        return store.max_marginal_relevance_search(question, k=top_k, fetch_k=top_k * 3)
    except Exception:
        try:
            scored = store.similarity_search_with_relevance_scores(question, k=top_k)
            return [d for d, _ in scored]
        except Exception:
            return store.similarity_search(question, k=top_k)


def retrieve_platform(question: str, k: int | None = None) -> list[Document]:
    settings = get_settings()
    service = get_rag_service()
    store = service.get_vectorstore()
    top_k = k or settings.retrieval_top_k
    try:
        scored = store.similarity_search_with_relevance_scores(question, k=top_k)
    except Exception:
        return store.similarity_search(question, k=top_k)
    threshold = settings.similarity_threshold
    kept = [doc for doc, score in scored if score is None or score >= threshold]
    return kept or [doc for doc, _ in scored]


def org_storage_dir(org_id: str, agent_id: str | None = None) -> Path:
    path = KB_UPLOAD_DIR / org_id
    if agent_id:
        path = path / agent_id
    path.mkdir(parents=True, exist_ok=True)
    return path


def avatar_storage_dir(org_id: str) -> Path:
    path = AVATAR_DIR / org_id
    path.mkdir(parents=True, exist_ok=True)
    return path


def verify_agent_embeddings(org_id: str, agent_id: str) -> dict[str, Any]:
    """Validate that Chroma vectors match DB chunk counts and embedding geometry is sane."""
    expected_dim = _probe_embedding_dim()
    store = get_agent_vectorstore(org_id, agent_id)
    collection = agent_collection_name(agent_id)
    total_vectors = int(store._collection.count() or 0)
    docs = [
        d
        for d in (saas_db.list_kb_documents(org_id, agent_id) or [])
        if d.get("status") == "ready"
    ]
    per_doc: list[dict[str, object]] = []
    issues: list[str] = []
    matched = 0
    for d in docs:
        doc_id = str(d.get("id") or "")
        chunk_count = int(d.get("chunk_count") or 0)
        emb_done = d.get("embeddings_done")
        vec_db = d.get("vectors_stored")
        try:
            got = store._collection.get(where={"doc_id": doc_id}, include=["embeddings", "documents"])
        except Exception as exc:
            issues.append(f"{d.get('filename')}: chroma get failed ({exc})")
            continue
        ids = got.get("ids") or []
        n_vec = len(ids)
        emb = got.get("embeddings")
        texts = got.get("documents") or []
        empty = sum(1 for t in texts if not (t or "").strip())
        dim = None
        norm = None
        if emb is not None and len(emb) > 0:
            try:
                import numpy as np

                e0 = np.asarray(emb[0], dtype=float)
                dim = int(e0.shape[-1])
                norm = float(np.linalg.norm(e0))
                if dim != expected_dim:
                    issues.append(f"{d.get('filename')}: embedding dim {dim} != {expected_dim}")
                if not bool(np.isfinite(e0).all()):
                    issues.append(f"{d.get('filename')}: non-finite embedding values")
                if abs(norm - 1.0) > 0.2:
                    issues.append(f"{d.get('filename')}: unusual L2 norm {norm:.4f}")
            except Exception as exc:
                issues.append(f"{d.get('filename')}: embedding inspect failed ({exc})")
        if chunk_count and n_vec != chunk_count:
            issues.append(
                f"{d.get('filename')}: chunk_count={chunk_count} but chroma_vectors={n_vec}"
            )
        else:
            matched += 1
        if empty:
            issues.append(f"{d.get('filename')}: {empty} empty chunk text(s)")
        per_doc.append(
            {
                "doc_id": doc_id,
                "filename": d.get("filename"),
                "chunk_count": chunk_count,
                "embeddings_done": emb_done,
                "vectors_stored_db": vec_db,
                "chroma_vectors": n_vec,
                "embed_dim": dim,
                "embed_norm": norm,
                "empty_chunks": empty,
                "ok": chunk_count == n_vec and empty == 0 and (dim in (None, expected_dim)),
            }
        )

    retrieval_ok = False
    top_score = None
    if total_vectors > 0:
        scored, _reason = retrieve_agent_scored(org_id, agent_id, "what is this document about", k=3)
        retrieval_ok = bool(scored)
        top_score = scored[0][1] if scored else None
        if not scored:
            issues.append("retrieval returned 0 hits for a basic overview query")

    ok = not issues and total_vectors > 0 and retrieval_ok
    return {
        "ok": ok,
        "embedding_model": get_settings().embedding_model,
        "expected_dim": expected_dim,
        "collection": collection,
        "total_vectors": total_vectors,
        "ready_docs": len(docs),
        "docs_matched": matched,
        "retrieval_ok": retrieval_ok,
        "top_score": top_score,
        "issues": issues,
        "documents": per_doc,
    }
