"""Coarse → fine retrieval: document → section → paragraph.""" from __future__ import annotations import logging from typing import TYPE_CHECKING from app.config import settings from app.models.schemas import SearchResult from app.retrieval.retriever import retrieve_document_level_context, retrieve_for_report if TYPE_CHECKING: pass logger = logging.getLogger(__name__) def _exemplar_doc_ids() -> list[str]: raw = (settings.rics_exemplar_document_ids or "").strip() if not raw: return [] return [x.strip() for x in raw.split(",") if x.strip()] def _unique_by_doc(scored: list[SearchResult], limit: int) -> list[str]: seen: set[str] = set() out: list[str] = [] for r in scored: if r.doc_id in seen: continue seen.add(r.doc_id) out.append(r.doc_id) if len(out) >= limit: break return out def retrieve_hierarchical_for_report( *, query: str, tenant_id: str, template_id: str, skeleton_excerpt: str, primary_document_id: str | None, reference_document_ids: list[str] | None, extra_doc_ids: list[str] | None = None, product_label: str | None = None, ) -> tuple[list[SearchResult], list[SearchResult], list[SearchResult]]: """Return ``(document_hits, section_hits, paragraph_pool)`` for generation. Paragraph hits should be reranked and truncated to ``rerank_top_n`` by the caller. When the index has no ``hierarchy_level=document`` rows (legacy ingest), routing falls back to prioritised flat retrieval for all three lists. """ from app.vectorstore.factory import get_vectorstore vs = get_vectorstore() refs = list(reference_document_ids or []) exemplars = _exemplar_doc_ids() secondary = list(dict.fromkeys(refs + exemplars)) extras = [x.strip() for x in (extra_doc_ids or []) if x and str(x).strip()] sk = (skeleton_excerpt or "").strip().replace("\n", " ")[:500] pl = (product_label or "RICS Home Survey Level 3 (Building Survey)").strip() doc_broad_query = ( f"{pl} report {template_id}. Whole document scope and narrative. {sk} {query[:400]}" ) k_doc = settings.hierarchical_k_document k_sec = settings.hierarchical_k_section k_pool = settings.hierarchical_k_paragraph_pool # ── Document-level ────────────────────────────────────────────────────── doc_hits = vs.search( doc_broad_query, tenant_id, k=max(k_doc * 25, 40), hierarchy_level="document", ) top_doc_ids = _unique_by_doc(doc_hits, max(k_doc, 6)) # Legacy index: no document-level rows if not doc_hits: flat = retrieve_for_report( query=query, tenant_id=tenant_id, primary_document_id=primary_document_id, secondary_document_ids=secondary, k=max(k_pool, 24), ) doc_ctx = retrieve_document_level_context( template_id=template_id, skeleton_excerpt=skeleton_excerpt, tenant_id=tenant_id, primary_document_id=primary_document_id, reference_document_ids=secondary, product_label=product_label, ) return doc_ctx, [], flat # Preferred doc universe: routing ∪ primary ∪ references ∪ runtime overlays universe: set[str] = set(top_doc_ids) if primary_document_id: universe.add(str(primary_document_id)) universe.update(secondary) universe.update(extras) allowed = frozenset(universe) doc_filtered = [r for r in doc_hits if r.doc_id in allowed][:k_doc] # Ensure document-tier rows exist for forced ids (e.g. runtime-edited sections) that did not rank globally. if extras: have_doc_tier = {r.doc_id for r in doc_filtered} for eid in extras: if eid in have_doc_tier: continue extra_rows = vs.search( doc_broad_query, tenant_id, k=12, hierarchy_level="document", doc_id_in=frozenset({eid}), ) if extra_rows: doc_filtered.append(extra_rows[0]) have_doc_tier.add(eid) doc_filtered = doc_filtered[:k_doc] # ── Section-level (within routed docs) ───────────────────────────────── sec_hits = vs.search( query, tenant_id, k=max(k_sec * 30, 60), hierarchy_level="section", doc_id_in=allowed, ) sec_filtered = sec_hits[:k_sec] # ── Paragraph pool ───────────────────────────────────────────────────── para_hits = vs.search( query, tenant_id, k=max(k_pool * 25, 80), hierarchy_level="paragraph", doc_id_in=allowed, ) if len(para_hits) < max(8, k_pool // 2): # Backfill from flat search (e.g. older paragraph rows without hierarchy tag) extra = retrieve_for_report( query=query, tenant_id=tenant_id, primary_document_id=primary_document_id, secondary_document_ids=secondary, k=k_pool, ) seen = {r.chunk_id for r in para_hits} for r in extra: if r.chunk_id not in seen: seen.add(r.chunk_id) para_hits.append(r) if len(para_hits) >= k_pool: break para_filtered = para_hits[:k_pool] logger.debug( "hierarchical RAG: doc=%d sec=%d para_pool=%d universe=%d", len(doc_filtered), len(sec_filtered), len(para_filtered), len(allowed), ) return doc_filtered, sec_filtered, para_filtered