study-buddy / app /rag /project_memory.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
9.11 kB
"""Chroma-backed project-specific research memory."""
from __future__ import annotations
import hashlib
import json
from datetime import datetime, timezone
from typing import Any, Sequence
from pydantic import BaseModel, Field
from app.observability.operation import observe_operation
from app.rag.chromadb_client import ChromaDBClient
from app.rag.models import MemoryContextItem
PROJECT_MEMORY_COLLECTION = "project_memory"
class ProjectMemoryRecord(BaseModel):
memory_id: str
project_id: str
kind: str
statement: str
evidence_ids: list[str] = Field(default_factory=list)
interaction_ids: list[str] = Field(default_factory=list)
attribution: str
confidence: float = Field(default=1.0, ge=0.0, le=1.0)
novelty_key: str = ""
status: str = "active"
contradicts: list[str] = Field(default_factory=list)
supersedes: list[str] = Field(default_factory=list)
observed_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc))
metadata: dict[str, Any] = Field(default_factory=dict)
def make_project_memory_id(project_id: str, kind: str, novelty_key: str, statement: str) -> str:
normalized = " ".join((novelty_key or statement).casefold().split())
digest = hashlib.sha256(f"{project_id}\0{kind}\0{normalized}".encode("utf-8")).hexdigest()
return f"pm_{digest}"
class ProjectMemoryStore:
"""Project-only personalization and research state backed by Chroma."""
def __init__(self, db: ChromaDBClient | None = None) -> None:
self.db = db or ChromaDBClient()
def upsert(self, records: ProjectMemoryRecord | Sequence[ProjectMemoryRecord]) -> int:
items = [records] if isinstance(records, ProjectMemoryRecord) else list(records)
if not items:
return 0
for record in items:
if not record.project_id:
raise ValueError("project memory requires project_id")
for start in range(0, len(items), 128):
batch = items[start : start + 128]
self.db.upsert(
PROJECT_MEMORY_COLLECTION,
documents=[_document(record) for record in batch],
metadatas=[_metadata(record) for record in batch],
ids=[record.memory_id for record in batch],
)
return len(items)
def search(
self,
project_id: str,
query: str,
limit: int = 8,
*,
kinds: Sequence[str] | None = None,
consumer: str | None = None,
) -> list[MemoryContextItem]:
if not project_id or not query.strip():
return []
clauses: list[dict[str, Any]] = [{"project_id": project_id}, {"status": "active"}]
if kinds:
clauses.append({"kind": {"$in": list(kinds)}})
with observe_operation(
"embedding.query",
subsystem="embedding",
consumer=consumer,
attributes={"collection_role": "project_memory"},
) as op:
embedding = self.db.embedder.embed([query])[0]
op.add_count("query_chars", len(query))
rows = self.db.query_raw(
PROJECT_MEMORY_COLLECTION,
embedding,
n_results=limit,
where={"$and": clauses},
consumer=consumer,
raise_on_failure=True,
)
return [_memory_hit(project_id, row) for row in rows]
def delete_project(self, project_id: str) -> None:
if project_id:
self.db.delete_where(PROJECT_MEMORY_COLLECTION, {"project_id": project_id})
def count_project(self, project_id: str) -> int:
return self.db.count_where(PROJECT_MEMORY_COLLECTION, {"project_id": project_id}) if project_id else 0
def list_project(self, project_id: str) -> list[MemoryContextItem]:
if not project_id:
return []
result = self.db.get_where(
PROJECT_MEMORY_COLLECTION,
{"$and": [{"project_id": project_id}, {"status": "active"}]},
include=["documents", "metadatas"],
)
ids = result.get("ids") or []
documents = result.get("documents") or []
metadatas = result.get("metadatas") or []
return [
_memory_hit(
project_id,
{
"id": memory_id,
"text": documents[index] if index < len(documents) else "",
"metadata": metadatas[index] if index < len(metadatas) else {},
"distance": None,
},
)
for index, memory_id in enumerate(ids)
]
def delete_document(self, project_id: str, document_id: str) -> None:
if project_id and document_id:
self.db.delete_where(
PROJECT_MEMORY_COLLECTION,
{"$and": [{"project_id": project_id}, {"document_id": document_id}]},
)
def delete_annotation_document(self, project_id: str, document_id: str) -> None:
if project_id and document_id:
self.db.delete_where(
PROJECT_MEMORY_COLLECTION,
{
"$and": [
{"project_id": project_id},
{"document_id": document_id},
{"kind": "annotation_note"},
]
},
)
def delete_by_evidence_ids(self, project_id: str, evidence_ids: Sequence[str]) -> int:
"""Remove memories grounded in evidence that is about to disappear."""
targets = set(evidence_ids)
if not project_id or not targets:
return 0
result = self.db.get_where(
PROJECT_MEMORY_COLLECTION,
{"project_id": project_id},
include=["metadatas"],
)
ids = result.get("ids") or []
metadatas = result.get("metadatas") or []
stale = [
str(memory_id)
for memory_id, metadata in zip(ids, metadatas)
if targets.intersection(_loads_list((metadata or {}).get("evidence_ids_json")))
]
if stale:
self.db.delete_ids(PROJECT_MEMORY_COLLECTION, stale)
return len(stale)
def _document(record: ProjectMemoryRecord) -> str:
return (
f"Project memory type: {record.kind}\n"
f"Attribution: {record.attribution}\n"
f"Statement: {record.statement}"
)
def _metadata(record: ProjectMemoryRecord) -> dict[str, Any]:
return {
"project_id": record.project_id,
"kind": record.kind,
"attribution": record.attribution,
"confidence": float(record.confidence),
"novelty_key": record.novelty_key,
"status": record.status,
"observed_at": record.observed_at.isoformat(),
"evidence_ids_json": json.dumps(record.evidence_ids, separators=(",", ":")),
"interaction_ids_json": json.dumps(record.interaction_ids, separators=(",", ":")),
"contradicts_json": json.dumps(record.contradicts, separators=(",", ":")),
"supersedes_json": json.dumps(record.supersedes, separators=(",", ":")),
"document_id": str(record.metadata.get("document_id") or ""),
"metadata_json": json.dumps(record.metadata, ensure_ascii=False, separators=(",", ":")),
}
def _memory_hit(project_id: str, row: dict[str, Any]) -> MemoryContextItem:
metadata = dict(row.get("metadata") or {})
document = str(row.get("text") or "")
marker = "Statement: "
statement = document.split(marker, 1)[1].strip() if marker in document else document.strip()
return MemoryContextItem(
memory_id=str(row["id"]),
source="project_memory",
statement=statement,
project_id=project_id,
kind=str(metadata.get("kind") or "project_observation"),
score=_distance_score(row.get("distance")),
evidence_ids=_loads_list(metadata.get("evidence_ids_json")),
observed_at=metadata.get("observed_at") or None,
metadata={
"attribution": metadata.get("attribution"),
"confidence": metadata.get("confidence"),
"novelty_key": metadata.get("novelty_key"),
"contradicts": _loads_list(metadata.get("contradicts_json")),
"supersedes": _loads_list(metadata.get("supersedes_json")),
"record": _loads_dict(metadata.get("metadata_json")),
},
)
def _distance_score(distance: Any) -> float:
if distance is None:
return 0.0
return 1.0 / (1.0 + max(0.0, float(distance)))
def _loads_list(value: Any) -> list[str]:
if not value:
return []
try:
parsed = json.loads(str(value))
return [str(item) for item in parsed] if isinstance(parsed, list) else []
except (TypeError, ValueError, json.JSONDecodeError):
return []
def _loads_dict(value: Any) -> dict[str, Any]:
if not value:
return {}
try:
parsed = json.loads(str(value))
return dict(parsed) if isinstance(parsed, dict) else {}
except (TypeError, ValueError, json.JSONDecodeError):
return {}