study-buddy / app /rag /vector_indexes.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
4.49 kB
"""Repository boundaries for Chroma-backed evidence and project memory."""
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Iterable, Sequence
from app.observability.operation import observe_operation
from app.rag.chromadb_client import ChromaDBClient
from app.rag.models import EvidenceType, EvidenceUnit, SourceKind
PAPER_EVIDENCE_COLLECTION = "paper_evidence"
@dataclass(frozen=True)
class VectorCandidate:
item_id: str
text: str
metadata: dict[str, Any]
distance: float | None
score: float
class PaperEvidenceIndex:
def __init__(self, db: ChromaDBClient | None = None) -> None:
self.db = db or ChromaDBClient()
def replace_document(
self,
project_id: str,
document_id: str,
units: Sequence[EvidenceUnit],
) -> int:
where = {"$and": [{"project_id": project_id}, {"document_id": document_id}]}
self.db.delete_where(PAPER_EVIDENCE_COLLECTION, where)
indexable = [
unit
for unit in units
if unit.index_text.strip() and unit.source_kind is not SourceKind.STUDENT_ANNOTATION
]
for batch in _batches(indexable, 128):
self.db.upsert(
PAPER_EVIDENCE_COLLECTION,
documents=[unit.index_text for unit in batch],
metadatas=[_evidence_metadata(unit) for unit in batch],
ids=[unit.evidence_id for unit in batch],
)
return len(indexable)
def search(
self,
project_id: str,
query: str,
limit: int = 30,
*,
document_ids: Sequence[str] | None = None,
evidence_types: set[EvidenceType] | None = None,
consumer: str | None = None,
) -> list[VectorCandidate]:
if not project_id or not query.strip():
return []
clauses: list[dict[str, Any]] = [{"project_id": project_id}]
if document_ids:
clauses.append({"document_id": {"$in": list(document_ids)}})
if evidence_types:
clauses.append({"element_type": {"$in": [item.value for item in evidence_types]}})
where = clauses[0] if len(clauses) == 1 else {"$and": clauses}
with observe_operation(
"embedding.query",
subsystem="embedding",
consumer=consumer,
attributes={"collection_role": "paper_evidence"},
) as op:
embedding = self.db.embedder.embed([query])[0]
op.add_count("query_chars", len(query))
rows = self.db.query_raw(
PAPER_EVIDENCE_COLLECTION,
embedding,
n_results=limit,
where=where,
consumer=consumer,
raise_on_failure=True,
)
return [_candidate(row) for row in rows]
def delete_document(self, project_id: str, document_id: str) -> None:
self.db.delete_where(
PAPER_EVIDENCE_COLLECTION,
{"$and": [{"project_id": project_id}, {"document_id": document_id}]},
)
def delete_project(self, project_id: str) -> None:
self.db.delete_where(PAPER_EVIDENCE_COLLECTION, {"project_id": project_id})
def _candidate(row: dict[str, Any]) -> VectorCandidate:
distance = row.get("distance")
return VectorCandidate(
item_id=str(row["id"]),
text=str(row.get("text") or ""),
metadata=dict(row.get("metadata") or {}),
distance=float(distance) if distance is not None else None,
score=_distance_score(distance),
)
def _distance_score(distance: Any) -> float:
if distance is None:
return 0.0
value = max(0.0, float(distance))
return 1.0 / (1.0 + value)
def _evidence_metadata(unit: EvidenceUnit) -> dict[str, Any]:
return {
"project_id": unit.project_id,
"document_id": unit.document_id,
"evidence_id": unit.evidence_id,
"element_type": unit.element_type.value,
"page_start": unit.page_start,
"page_end": unit.page_end,
"parent_id": unit.parent_id or "",
"section_path": " > ".join(unit.section_path),
"source_kind": unit.source_kind.value,
"bbox_json": json.dumps(unit.bbox_norm.rounded(), separators=(",", ":")) if unit.bbox_norm else "",
}
def _batches(items: Sequence[Any], size: int) -> Iterable[Sequence[Any]]:
for start in range(0, len(items), size):
yield items[start : start + size]