study-buddy / app /rag /chromadb_client.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
14.5 kB
import hashlib
from pathlib import Path
from typing import Any, Dict, List, Optional
import chromadb
from chromadb.errors import NotFoundError
from app.observability import retrieval as _obs_retrieval
from app.observability.contracts import RetrievalOutcome
from app.observability.operation import observe_operation
from app.rag.embeddings import Embedder
from app.rag.pipeline_version import get_pipeline_version
_FILE_INDEX_COLLECTION = "file-index"
_COLLECTION_ROLES = {
"paper_evidence": "paper_evidence",
"project_memory": "project_memory",
}
def _collection_role(collection: str) -> str:
"""Return the finite role name safe to attach to retrieval spans."""
return _COLLECTION_ROLES.get(collection, "other")
def _safe_count(col: Any) -> Optional[int]:
"""Read a collection's size without ever raising into product code."""
try:
return col.count()
except Exception:
return None
class ChromaDBClient:
"""In-memory ChromaDB wrapper for the web demo deployment.
State is persistent across process restarts. Files are deduplicated by
SHA-256 content hash.
"""
def __init__(
self,
embedder: Optional[Embedder] = None,
*,
path: str | Path | None = None,
client: Any | None = None,
) -> None:
import os
if client is not None:
self._client = client
else:
db_path = Path(path).expanduser() if path is not None else Path(os.path.expanduser("~/.studybuddy/chroma"))
db_path.mkdir(parents=True, exist_ok=True)
self._client = chromadb.PersistentClient(path=str(db_path))
self.embedder = embedder or Embedder()
# ------------------------------------------------------------------ #
# Internal #
# ------------------------------------------------------------------ #
def _get_or_create(self, name: str):
return self._client.get_or_create_collection(name)
# ------------------------------------------------------------------ #
# File dedup #
# ------------------------------------------------------------------ #
@staticmethod
def file_hash(content: bytes) -> str:
return hashlib.sha256(content).hexdigest()
def is_indexed(self, content_hash: str) -> bool:
try:
col = self._client.get_or_create_collection(_FILE_INDEX_COLLECTION)
r = col.get(ids=[content_hash])
return len(r["ids"]) > 0
except Exception:
return False
def mark_indexed(self, content_hash: str, filename: str) -> None:
col = self._client.get_or_create_collection(_FILE_INDEX_COLLECTION)
col.upsert(
ids=[content_hash],
documents=[filename],
metadatas=[{"filename": filename}],
)
# ------------------------------------------------------------------ #
# Public API #
# ------------------------------------------------------------------ #
def upsert(
self,
collection: str,
documents: List[str],
metadatas: List[Dict[str, Any]],
ids: List[str],
) -> None:
col = self._get_or_create(collection)
with observe_operation("embedding.batch", subsystem="embedding") as op:
op.add_count("embedding_batch_size", len(documents))
embeddings = self.embedder.embed(documents)
with observe_operation("chroma.upsert", subsystem="chroma") as op:
before = _safe_count(col)
if before is not None:
op.set("collection_count_before", before)
col.upsert(documents=documents, embeddings=embeddings, metadatas=metadatas, ids=ids)
after = _safe_count(col)
if after is not None:
op.set("collection_count_after", after)
def query_observed(
self,
collection: str,
query_embedding: List[float],
n_results: int = 5,
where: Optional[Dict[str, Any]] = None,
*,
consumer: Optional[str] = None,
pipeline_version: Optional[str] = None,
) -> RetrievalOutcome:
"""Vector search that preserves the *cause* of an empty result.
Behaviourally identical to :meth:`query` (same clamping, same
fallbacks), but returns a :class:`RetrievalOutcome` so a caller can
tell a genuinely-empty successful search (``success_empty``) apart from
a collection-lookup / vector-search failure that fell back to ``[]``
(``error_fallback``). :meth:`query` delegates here and unwraps
``.candidates`` so every existing caller stays byte-for-byte compatible.
``pipeline_version`` (default ``None``) gates the over-fetch + dedup +
truncate mechanism -- see ``app.rag.pipeline_version``. ``None`` (every
ordinary product call site) and the explicit ``"rag-naive-v1"`` both
resolve to the exact same frozen, unmodified behavior this method has
always had: request exactly ``n_results``, no over-fetch, no dedup.
Only an explicit ``"rag-dedup-v2"`` (driven by the benchmark runner)
changes anything.
"""
pv = get_pipeline_version(pipeline_version)
requested_n = n_results # what the caller actually wants back, unchanged by over-fetch
with observe_operation("chroma.collection_lookup", subsystem="chroma", consumer=consumer) as op:
try:
col = self._client.get_collection(collection)
except Exception as exc:
op.mark_error(_obs_retrieval.COLLECTION_UNAVAILABLE)
return RetrievalOutcome.failure(_obs_retrieval.COLLECTION_UNAVAILABLE)
# Over-fetch a larger raw pool when dedup is enabled, so removing
# near-duplicates below can still backfill up to `requested_n` from
# alternate candidates instead of only ever shrinking the naive
# top-k in place. `fetch_n == requested_n` whenever dedup is
# disabled (v1's overfetch_factor is pinned to 1), which is what
# keeps v1 byte-identical to its pre-Task-8 behavior.
fetch_n = requested_n * pv.overfetch_factor
# Clamp fetch_n to actual collection size (a count failure is
# non-fatal here, exactly as before). `requested_top_k` is always the
# caller's original ask, never the inflated over-fetch amount, so
# this attribute means the same thing across both pipeline versions.
with observe_operation("chroma.collection_count", subsystem="chroma", consumer=consumer) as op:
try:
actual = col.count()
op.set("collection_size", actual)
op.set("requested_top_k", requested_n)
fetch_n = min(fetch_n, actual) if actual > 0 else 1
except Exception:
pass
kwargs: Dict[str, Any] = {
"query_embeddings": [query_embedding],
"n_results": fetch_n,
}
if where:
kwargs["where"] = where
with observe_operation("chroma.vector_search", subsystem="chroma", consumer=consumer) as op:
try:
results = col.query(**kwargs)
except Exception as exc:
op.mark_error(_obs_retrieval.VECTOR_SEARCH_FAILED)
return RetrievalOutcome.failure(_obs_retrieval.VECTOR_SEARCH_FAILED)
with observe_operation("rag.result_prepare", subsystem="retrieval", consumer=consumer) as op:
out = []
for doc, meta in zip(results["documents"][0], results["metadatas"][0]):
out.append({"text": doc, **meta})
# Pool-level removed count, captured *before* truncation to
# `requested_n` -- this is direct evidence of what the dedup
# mechanism actually did to the raw over-fetched pool. It is
# deliberately distinct from `duplicate_candidate_ratio` below,
# which is computed on the final, already-deduplicated/truncated
# list `candidate_stats()` receives: a post-dedup set is ~0.0 by
# construction whether or not anything was removed, so it alone
# cannot tell "removed nothing" apart from "removed something and
# truncated it away". `duplicates_removed_from_pool` closes that
# gap. Always 0 for v1 (dedup disabled -- there is no pool step to
# remove anything from), which is itself the correct, honest value.
duplicates_removed_from_pool = 0
if pv.dedup_enabled:
dup_indices = _obs_retrieval.find_duplicate_indices(out)
duplicates_removed_from_pool = len(dup_indices)
out = [c for i, c in enumerate(out) if i not in dup_indices]
out = out[:requested_n]
op.set("duplicates_removed_from_pool", duplicates_removed_from_pool)
for key, value in _obs_retrieval.candidate_stats(out).items():
op.set(key, value)
outcome = RetrievalOutcome.success(out)
return outcome
def query(
self,
collection: str,
query_embedding: List[float],
n_results: int = 5,
where: Optional[Dict[str, Any]] = None,
*,
pipeline_version: Optional[str] = None,
) -> List[Dict[str, Any]]:
return self.query_observed(
collection,
query_embedding,
n_results=n_results,
where=where,
pipeline_version=pipeline_version,
).candidates
def query_raw(
self,
collection: str,
query_embedding: List[float],
n_results: int = 20,
where: Optional[Dict[str, Any]] = None,
*,
consumer: Optional[str] = None,
raise_on_failure: bool = False,
) -> List[Dict[str, Any]]:
"""Return stable IDs and distances for repository-level retrieval.
Product consumers should use ``PaperEvidenceService`` or the project
memory repository. This low-level method exists so rank fusion retains
Chroma's identity and score instead of flattening results into text.
``raise_on_failure`` lets the structured source-retrieval boundary
distinguish an unavailable dense index from a genuinely empty one;
legacy callers retain the established empty-list fallback.
"""
collection_role = _collection_role(collection)
with observe_operation(
"chroma.collection_lookup",
subsystem="chroma",
consumer=consumer,
attributes={"collection_role": collection_role},
) as op:
try:
col = self._client.get_collection(collection)
except NotFoundError:
return []
except Exception:
op.mark_error(_obs_retrieval.COLLECTION_UNAVAILABLE)
if raise_on_failure:
raise
return []
try:
actual = col.count()
except Exception:
actual = n_results
if actual <= 0:
return []
kwargs: Dict[str, Any] = {
"query_embeddings": [query_embedding],
"n_results": min(max(1, n_results), actual),
"include": ["documents", "metadatas", "distances"],
}
if where:
kwargs["where"] = where
with observe_operation(
"chroma.vector_search",
subsystem="chroma",
consumer=consumer,
attributes={"collection_role": collection_role},
) as op:
try:
result = col.query(**kwargs)
except Exception:
op.mark_error(_obs_retrieval.VECTOR_SEARCH_FAILED)
if raise_on_failure:
raise
return []
ids = (result.get("ids") or [[]])[0]
documents = (result.get("documents") or [[]])[0]
metadatas = (result.get("metadatas") or [[]])[0]
distances = (result.get("distances") or [[]])[0]
rows: List[Dict[str, Any]] = []
for index, item_id in enumerate(ids):
rows.append(
{
"id": item_id,
"text": documents[index] if index < len(documents) else "",
"metadata": metadatas[index] if index < len(metadatas) else {},
"distance": distances[index] if index < len(distances) else None,
}
)
return rows
def delete_where(self, collection: str, where: Dict[str, Any]) -> None:
"""Delete matching rows, treating only an absent collection as empty.
Derived-state deletion is a correctness operation. Authentication,
storage, or filter failures must reach the caller so source files are
not removed while stale vectors remain searchable.
"""
try:
col = self._client.get_collection(collection)
except NotFoundError:
return
col.delete(where=where)
def get_where(
self,
collection: str,
where: Dict[str, Any],
*,
include: List[str] | None = None,
) -> Dict[str, Any]:
try:
col = self._client.get_collection(collection)
except NotFoundError:
return {"ids": [], "metadatas": [], "documents": []}
return col.get(where=where, include=include or ["metadatas"])
def delete_ids(self, collection: str, ids: List[str]) -> None:
if not ids:
return
try:
col = self._client.get_collection(collection)
except NotFoundError:
return
col.delete(ids=ids)
def collection_count(self, collection: str) -> int:
try:
return self._client.get_collection(collection).count()
except Exception:
return 0
def count_where(self, collection: str, where: Dict[str, Any]) -> int:
try:
result = self._client.get_collection(collection).get(where=where, include=[])
return len(result.get("ids") or [])
except Exception:
return 0
def delete_collection(self, name: str) -> None:
try:
self._client.delete_collection(name)
except Exception:
pass