from __future__ import annotations
import json
import logging
import os
from collections import Counter
from dataclasses import dataclass, field
from typing import Any, Dict, Iterable, List, Optional, Tuple
from langchain_chroma import Chroma
from langchain_community.retrievers import BM25Retriever
from langchain_core.documents import Document
from langchain_ollama import OllamaEmbeddings
from .multi_vector_fusion import reciprocal_rank_fusion
[docs]
LOGGER = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# JSONL helpers
# ---------------------------------------------------------------------------
[docs]
def _iter_jsonl(jsonl_path: str) -> Iterable[Dict[str, Any]]:
with open(jsonl_path, "r", encoding="utf-8") as f:
for lineno, line in enumerate(f, start=1):
line = line.strip()
if not line:
continue
try:
obj = json.loads(line)
except json.JSONDecodeError as exc:
LOGGER.warning("_iter_jsonl: parse error in %s line %d: %s", jsonl_path, lineno, exc)
continue
if isinstance(obj, dict):
yield obj
[docs]
def _collapse_ws(s: Optional[str]) -> str:
return " ".join((s or "").split())
# ---------------------------------------------------------------------------
# Processed record helpers
# ---------------------------------------------------------------------------
[docs]
def _looks_like_processed_text_record(obj: Dict[str, Any]) -> bool:
return (
isinstance(obj, dict)
and all(k in obj for k in ("record_id", "doc_id", "doc_type", "chunk_index", "embedding_text", "metadata", "provenance"))
and isinstance(obj.get("metadata"), dict)
and isinstance(obj.get("provenance"), dict)
and isinstance(obj.get("record_id"), str)
and bool(obj.get("record_id", "").strip())
)
[docs]
def _stable_record_id_from_doc(doc: Document) -> Optional[str]:
metadata = dict(getattr(doc, "metadata", {}) or {})
for key in ("record_id", "id", "_id"):
value = metadata.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
for attr in ("id", "record_id"):
value = getattr(doc, attr, None)
if isinstance(value, str) and value.strip():
return value.strip()
return None
[docs]
LIST_FIELDS = {
"equipment_ids",
"system_names",
"component_names",
"mechanisms",
"failure_outcomes",
"maintenance_actions",
"surveillance_actions",
"tools_methods",
"properties_or_limits",
"doc_refs",
"alarm_ids",
}
[docs]
def _doc_matches_component_ids(doc: Document, wanted: set) -> bool:
"""Return True if *doc* is associated with at least one of the *wanted* component IDs.
Checks:
1. The scalar ``primary_component_id`` metadata field (index-friendly path).
2. The scalar ``component_id`` metadata field (legacy single-component path).
3. The ``component_ids`` field. In practice this is a JSON-encoded string on every
path, because metadata is passed through :func:`_sanitize_meta` (which stringifies
lists) before it reaches either Chroma or the in-memory BM25 corpus. The plain-list
branch is kept only as a defensive fallback for externally constructed documents.
"""
meta = doc.metadata or {}
primary = meta.get("primary_component_id")
if primary and primary in wanted:
return True
scalar = meta.get("component_id")
if scalar and scalar in wanted:
return True
raw = meta.get("component_ids")
if isinstance(raw, str):
try:
ids = json.loads(raw)
if isinstance(ids, list) and any(c in wanted for c in ids):
return True
except (json.JSONDecodeError, TypeError):
pass
elif isinstance(raw, (list, tuple)):
if any(c in wanted for c in raw):
return True
return False
[docs]
def _doc_matches_filter_sane(doc: Document, filter_sane: Dict[str, Any]) -> bool:
"""Apply the dense-path Chroma filter semantics to a BM25 hit's metadata.
The dense path builds ``{k: {"$in": vals}}`` for list-valued filters and ``{k: {"$eq": v}}``
for scalars. This mirrors that membership-vs-equality logic for the in-memory BM25
post-filter so the two retrieval views agree. Previously BM25 used ``==`` even for
list-valued filters, which never matched a scalar metadata value and silently dropped
every BM25 hit — collapsing hybrid retrieval to dense-only whenever a list filter was set.
"""
meta = doc.metadata or {}
for k, v in filter_sane.items():
actual = meta.get(k)
if isinstance(v, (list, tuple, set)):
if actual not in set(v):
return False
elif actual != v:
return False
return True
[docs]
def _record_component_ids(metadata: Dict[str, Any]) -> List[str]:
out: List[str] = []
scalar = metadata.get("component_id")
if isinstance(scalar, str) and scalar.strip():
out.append(scalar.strip())
raw = metadata.get("component_ids")
if isinstance(raw, str):
try:
decoded = json.loads(raw)
if isinstance(decoded, list):
out.extend(str(x).strip() for x in decoded if str(x).strip())
except (json.JSONDecodeError, TypeError, ValueError):
pass
elif isinstance(raw, (list, tuple, set)):
out.extend(str(x).strip() for x in raw if str(x).strip())
seen = set()
deduped: List[str] = []
for cid in out:
if cid not in seen:
seen.add(cid)
deduped.append(cid)
return deduped
[docs]
def to_chroma_payload(record: Dict[str, Any]) -> Dict[str, Any]:
"""Convert a processed_text_record into the ``(id, document, metadata)`` triple upserted
into Chroma.
The vector id is the record's ``record_id``, the embedded/indexed document text is the
whitespace-collapsed ``embedding_text``, and metadata comes from :func:`build_chroma_metadata`.
"""
return {
"id": str(record["record_id"]),
"document": _collapse_ws(record.get("embedding_text") or ""),
"metadata": build_chroma_metadata(record),
}
[docs]
def collection_name_for_doc_type(doc_type: str, prefix: str = "processed") -> str:
"""Return the Chroma collection name for a document type, e.g. ``processed_CR``.
The doc type is upper-cased and sanitized (``/`` and spaces become ``_``) so it is a
valid, stable collection identifier; a falsy doc type maps to ``OTHER``.
"""
dt = (doc_type or "OTHER").strip().upper().replace("/", "_").replace(" ", "_")
return f"{prefix}_{dt}"
@dataclass
[docs]
class _CollectionState:
[docs]
bm25_docs: List[Document] = field(default_factory=list)
[docs]
bm25: Optional[BM25Retriever] = None
[docs]
class ChromaRecordStore:
"""
Stage-6 store for canonical processed_text_record objects.
One Chroma collection is used per document class (CR, MR, SOP, ECA, ...).
Every indexed vector corresponds to one validated processed_text_record.
"""
def __init__(
self,
persist_directory: str,
*,
embed_model: Optional[str] = None,
ollama_base_url: Optional[str] = None,
collection_prefix: str = "processed",
bm25_k: int = 20,
) -> None:
"""Initialise the per-doc-type Chroma store.
Args:
persist_directory: On-disk directory for Chroma's persistent collections
(created if missing).
embed_model: Ollama embedding model name; defaults to ``$OLLAMA_EMBED_MODEL``
or ``mxbai-embed-large:335m``.
ollama_base_url: Ollama server URL; defaults to ``$OLLAMA_BASE_URL`` or
``http://localhost:11434``.
collection_prefix: Prefix for generated collection names (see
:func:`collection_name_for_doc_type`).
bm25_k: Default fan-out for the in-memory BM25 retriever built per collection.
Note:
The embedding model name is baked into each collection name, so switching models
transparently targets a distinct set of collections rather than mixing vector spaces.
"""
[docs]
self.persist_directory = persist_directory
os.makedirs(self.persist_directory, exist_ok=True)
[docs]
self.embed_model = embed_model or os.environ.get("OLLAMA_EMBED_MODEL", "mxbai-embed-large:335m")
[docs]
self.ollama_base_url = ollama_base_url or os.environ.get("OLLAMA_BASE_URL", "http://localhost:11434")
[docs]
self.collection_prefix = collection_prefix
[docs]
self.embedder = OllamaEmbeddings(base_url=self.ollama_base_url, model=self.embed_model)
[docs]
self._states: Dict[str, _CollectionState] = {}
[docs]
def _collection_name(self, doc_type: str) -> str:
safe_model = self.embed_model.replace(":", "_").replace("/", "_")
base = collection_name_for_doc_type(doc_type=doc_type, prefix=self.collection_prefix)
return f"{base}_{safe_model}"
[docs]
def get_or_create_collection(self, doc_type: str, collection_name: Optional[str] = None) -> _CollectionState:
cname = collection_name or self._collection_name(doc_type)
if cname in self._states:
return self._states[cname]
vs = Chroma(
collection_name=cname,
embedding_function=self.embedder,
persist_directory=self.persist_directory,
)
state = _CollectionState(vectorstore=vs)
self._states[cname] = state
LOGGER.info("ChromaRecordStore: opened/created collection '%s'.", cname)
return state
[docs]
def load_collection(self, doc_type: str, collection_name: Optional[str] = None) -> _CollectionState:
"""Open a persisted collection for querying and return its :class:`_CollectionState`.
The dense (vector) side is fully backed by Chroma's on-disk index, so it works
regardless of which process performed the original ingest.
BM25, however, is an *in-memory* corpus built only from documents upserted during the
current process (see :meth:`upsert_records`). When a collection is loaded fresh from
disk without a matching in-process ingest, ``state.bm25_docs`` is empty and hybrid
retrieval degrades to **dense-only** — ``hybrid_weight`` then has no effect. This is a
deliberate limitation of the current design (the BM25 corpus is not persisted); callers
needing hybrid retrieval in a query-only process must re-ingest the source JSONL first.
A warning is emitted here to make the degraded mode explicit.
"""
state = self.get_or_create_collection(doc_type=doc_type, collection_name=collection_name)
if state.bm25 is None and not state.bm25_docs:
LOGGER.warning(
"Collection '%s' loaded from disk without an in-process ingest — BM25 corpus is "
"empty, so retrieval will be dense-only (hybrid_weight has no effect) until the "
"source records are re-ingested in this process.",
collection_name or self._collection_name(doc_type),
)
return state
[docs]
def _get_or_build_bm25(self, state: _CollectionState) -> Optional[BM25Retriever]:
if state.bm25 is not None:
return state.bm25
if not state.bm25_docs:
return None
state.bm25 = BM25Retriever.from_documents(state.bm25_docs, k=self.bm25_k)
return state.bm25
[docs]
def upsert_records(
self,
records: Iterable[Dict[str, Any]],
*,
doc_type: Optional[str] = None,
collection_name: Optional[str] = None,
) -> int:
"""Upsert a batch of processed_text_records into the collection for their doc type.
All records in a call must share one doc type (a mixed-type batch raises
``ValueError``); the type is taken from ``doc_type`` or inferred from the first record.
Malformed records and records with empty embedding text are skipped. Each surviving
record is upserted into Chroma by ``record_id`` (dense side) and added to the
collection's in-memory BM25 corpus (sparse side), so a subsequent query in the same
process gets true hybrid retrieval.
Returns:
The number of records actually upserted.
"""
records = list(records)
if not records:
return 0
inferred_doc_type = (doc_type or str(records[0].get("doc_type") or "OTHER")).strip().upper()
bad_doc_types: List[str] = []
for rec in records:
rec_dt = str(rec.get("doc_type") or inferred_doc_type).strip().upper()
if rec_dt != inferred_doc_type:
bad_doc_types.append(str(rec.get("record_id") or "UNKNOWN"))
if bad_doc_types:
raise ValueError(
f"upsert_records received mixed doc_type batch for collection doc_type={inferred_doc_type}. "
f"Offending record_ids={bad_doc_types[:10]}"
)
state = self.get_or_create_collection(doc_type=inferred_doc_type, collection_name=collection_name)
vs = state.vectorstore
ids: List[str] = []
docs: List[str] = []
metas: List[Dict[str, Any]] = []
bm25_docs: List[Document] = []
for rec in records:
if not _looks_like_processed_text_record(rec):
LOGGER.warning("Skipping malformed processed_text_record during Chroma upsert: keys=%s", sorted(rec.keys()))
continue
payload = to_chroma_payload(rec)
if not payload["document"]:
continue
ids.append(payload["id"])
docs.append(payload["document"])
metas.append(payload["metadata"])
bm25_docs.append(Document(page_content=payload["document"], metadata=payload["metadata"]))
if not ids:
return 0
# Use the public LangChain vectorstore API; it embeds the documents and upserts by id.
vs.add_texts(texts=docs, metadatas=metas, ids=ids)
# Keep a fresh BM25 corpus for collections ingested in-process.
existing_by_id = {d.metadata.get("record_id"): d for d in state.bm25_docs}
for d in bm25_docs:
rid = d.metadata.get("record_id")
if rid:
existing_by_id[rid] = d
state.bm25_docs = list(existing_by_id.values())
state.bm25 = None
LOGGER.info("ChromaRecordStore: upserted %d records into '%s'.", len(ids), vs._collection.name)
return len(ids)
[docs]
def upsert_jsonl(
self,
jsonl_path: str,
*,
doc_type_override: Optional[str] = None,
) -> Dict[str, int]:
"""Ingest a JSONL file of processed_text_records, one Chroma collection per doc type.
Records are read from ``jsonl_path`` (each line may be a bare record or wrapped under a
``processed_text_record`` key), grouped by doc type (or forced to ``doc_type_override``),
and upserted via :meth:`upsert_records`.
Returns:
Mapping of ``doc_type -> number of records upserted``.
"""
grouped: Dict[str, List[Dict[str, Any]]] = {}
for obj in _iter_jsonl(jsonl_path):
rec = extract_processed_text_record(obj)
if not rec:
continue
doc_type = (doc_type_override or rec.get("doc_type") or "OTHER").strip().upper()
grouped.setdefault(doc_type, []).append(rec)
counts: Dict[str, int] = {}
for doc_type, records in grouped.items():
counts[doc_type] = self.upsert_records(records, doc_type=doc_type)
return counts
[docs]
def query_doc_type(
self,
doc_type: str,
query_text: str,
*,
top_k: int = 8,
filter_meta: Optional[Dict[str, Any]] = None,
collection_name: Optional[str] = None,
hybrid_weight: float = 0.5,
) -> List[Document]:
"""Hybrid (dense + BM25) retrieval over a single doc-type collection.
Runs a dense vector search and a BM25 search, then fuses the two ranked lists with
Reciprocal Rank Fusion weighted by ``hybrid_weight`` (dense) / ``1 - hybrid_weight``
(BM25). The fused ``_score`` is written back onto each returned document's metadata.
Filtering:
``filter_meta`` is normalized (see :func:`_normalize_filter_meta`) into scalar
Chroma ``$eq``/``$in`` clauses. ``component_ids`` is handled specially: it becomes
an index-level ``primary_component_id`` ``$in`` filter, with a legacy post-filter
fallback for older records that predate ``primary_component_id`` (see
:func:`_doc_matches_component_ids`).
BM25 availability:
BM25 only contributes when this collection was ingested in the current process; on
a disk-loaded collection retrieval is dense-only (see :meth:`load_collection`).
``_bm25_available`` is recorded on each returned document's metadata.
Args:
doc_type: Document type selecting the collection.
query_text: Natural-language query.
top_k: Maximum number of fused results to return.
filter_meta: Optional high-level metadata filters.
collection_name: Explicit collection override (else derived from ``doc_type``).
hybrid_weight: Dense-vs-BM25 blend in [0, 1]; 1.0 is dense-only, 0.0 is BM25-only.
Returns:
Up to ``top_k`` LangChain ``Document`` objects ordered by fused score.
Raises:
ValueError: If the target collection has not been initialised via
:meth:`upsert_jsonl` / :meth:`upsert_records` / :meth:`load_collection`.
"""
cname = collection_name or self._collection_name(doc_type)
if cname not in self._states:
raise ValueError(f"Collection '{cname}' is not initialised. Call upsert_jsonl() or load_collection() first.")
state = self._states[cname]
vs = state.vectorstore
filter_sane = _normalize_filter_meta(filter_meta)
# component_ids is translated to index-level primary_component_id filtering.
# Legacy records without primary_component_id are optionally recovered through
# a compatibility post-filter pass against component_ids/component_id metadata.
wanted_component_ids: set = set((filter_meta or {}).get("component_ids") or [])
if not filter_sane:
chroma_where = None
else:
clauses = []
for k, v in filter_sane.items():
if isinstance(v, (list, tuple, set)):
vals = [x for x in v if x is not None]
if not vals:
continue
clauses.append({k: {"$in": list(vals)}})
else:
clauses.append({k: {"$eq": v}})
if not clauses:
chroma_where = None
elif len(clauses) == 1:
chroma_where = clauses[0]
else:
chroma_where = {"$and": clauses}
dense_hits: List[Dict[str, Any]] = []
fetch_k = top_k * 2
fallback_fetch_k = top_k * 4
indexed_component_filter_active = bool(wanted_component_ids)
if indexed_component_filter_active:
comp_clause = {"primary_component_id": {"$in": sorted(wanted_component_ids)}}
if chroma_where is None:
chroma_where = comp_clause
elif "$and" in chroma_where and isinstance(chroma_where.get("$and"), list):
chroma_where = {"$and": list(chroma_where["$and"]) + [comp_clause]}
else:
chroma_where = {"$and": [chroma_where, comp_clause]}
try:
scored = vs.similarity_search_with_score(query_text, k=fetch_k, filter=chroma_where)
for doc, score in scored:
doc.metadata = dict(doc.metadata or {})
doc.metadata["_score"] = float(score)
# Preserve raw vector similarity before RRF fusion overwrites _score.
# Downstream consumers (e.g. evidence assessment) use _vector_score
# as a semantic relevance signal that survives RRF re-ranking.
doc.metadata["_vector_score"] = float(score)
rid = _stable_record_id_from_doc(doc)
if not rid:
LOGGER.warning("Dense retrieval returned hit with no stable record_id; skipping.")
continue
if wanted_component_ids:
doc.metadata["_component_filter_strategy"] = "index_filter"
dense_hits.append({
"record_id": rid,
"score": float(score),
"document": doc,
"metadata": doc.metadata,
})
except Exception as exc:
LOGGER.warning("Dense retrieval failed for doc_type '%s': %s", doc_type, exc)
if wanted_component_ids and len(dense_hits) < top_k:
# Compatibility path for older ingested records that do not yet carry
# primary_component_id metadata.
try:
legacy_scored = vs.similarity_search_with_score(query_text, k=fallback_fetch_k, filter=filter_sane or None)
seen = {h.get("record_id") for h in dense_hits}
legacy_added = 0
for doc, score in legacy_scored:
rid = _stable_record_id_from_doc(doc)
if not rid or rid in seen:
continue
doc.metadata = dict(doc.metadata or {})
if doc.metadata.get("primary_component_id"):
continue
if not _doc_matches_component_ids(doc, wanted_component_ids):
continue
doc.metadata["_score"] = float(score)
doc.metadata["_vector_score"] = float(score)
doc.metadata["_component_filter_strategy"] = "legacy_post_filter"
dense_hits.append(
{
"record_id": rid,
"score": float(score),
"document": doc,
"metadata": doc.metadata,
}
)
seen.add(rid)
legacy_added += 1
if legacy_added:
LOGGER.warning(
"ChromaRecordStore: component filter used index-level primary_component_id "
"plus legacy post-filter fallback for %d records lacking primary_component_id. "
"Re-ingest collections to remove fallback dependence.",
legacy_added,
)
except Exception as exc:
LOGGER.warning("Legacy component fallback retrieval failed for doc_type '%s': %s", doc_type, exc)
bm25_hits: List[Dict[str, Any]] = []
bm25 = self._get_or_build_bm25(state)
bm25_available = bm25 is not None
if bm25_available:
try:
bm25.k = fetch_k
docs = bm25.invoke(query_text)
if filter_sane or wanted_component_ids:
docs = [
d for d in docs
if _doc_matches_filter_sane(d, filter_sane)
and (not wanted_component_ids or _doc_matches_component_ids(d, wanted_component_ids))
]
for doc in docs:
if wanted_component_ids and "_component_filter_strategy" not in (doc.metadata or {}):
doc.metadata = dict(doc.metadata or {})
if doc.metadata.get("primary_component_id"):
doc.metadata["_component_filter_strategy"] = "index_filter"
else:
doc.metadata["_component_filter_strategy"] = "legacy_post_filter"
rid = _stable_record_id_from_doc(doc)
if not rid:
LOGGER.warning("BM25 retrieval returned hit with no stable record_id; skipping.")
continue
bm25_hits.append({
"record_id": rid,
"score": None,
"document": doc,
"metadata": doc.metadata,
})
except Exception as exc:
bm25_available = False
LOGGER.warning("BM25 retrieval failed for doc_type '%s': %s", doc_type, exc)
else:
LOGGER.warning(
"BM25 unavailable for collection '%s' (loaded from disk without in-process ingest); "
"falling back to dense-only retrieval. hybrid_weight=%.2f has no effect.",
cname,
hybrid_weight,
)
fused = reciprocal_rank_fusion(
{"dense": dense_hits, "bm25": bm25_hits},
k=top_k,
view_weights={"dense": hybrid_weight, "bm25": 1.0 - hybrid_weight},
)
out: List[Document] = []
for entry in fused:
doc = entry.get("document")
if doc is None:
continue
doc.metadata = dict(doc.metadata or {})
doc.metadata["_score"] = entry["score"]
doc.metadata["_bm25_available"] = bm25_available
out.append(doc)
return out