Spaces:
Sleeping
Sleeping
| """ | |
| ChromaDB vector store for persistent embedding storage. | |
| Collections are named {session_id}__{model_name} so the same conversation | |
| can be embedded with different models and compared. | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import threading | |
| from typing import List, Optional | |
| import numpy as np | |
| import chromadb | |
| logger = logging.getLogger(__name__) | |
| # ChromaDB's Rust backend is not thread-safe for concurrent initialization. | |
| # Serialize all PersistentClient creation to avoid segfaults / AttributeErrors. | |
| _chromadb_init_lock = threading.Lock() | |
| class VectorStore: | |
| """ChromaDB wrapper for storing and retrieving embeddings. | |
| Args: | |
| persist_dir: Directory for ChromaDB persistent storage. | |
| """ | |
| def __init__(self, persist_dir: str): | |
| with _chromadb_init_lock: | |
| self._client = chromadb.PersistentClient(path=persist_dir) | |
| def _collection_name(session_id: str, model_name: str) -> str: | |
| """Build a collection name from session ID and model name. | |
| ChromaDB collection names must be 3-63 chars, start/end with | |
| alphanumeric, and contain only alphanumerics, underscores, hyphens. | |
| """ | |
| raw = f"{session_id}__{model_name}" | |
| sanitized = "".join(c if c.isalnum() or c in ("_", "-") else "_" for c in raw) | |
| if len(sanitized) < 3: | |
| sanitized = sanitized + "___" | |
| return sanitized[:63] | |
| def store_embeddings( | |
| self, | |
| session_id: str, | |
| model_name: str, | |
| texts: List[str], | |
| embeddings: np.ndarray, | |
| metadatas: Optional[List[dict]] = None, | |
| ): | |
| """Store embeddings for a session+model pair. | |
| Args: | |
| session_id: Session identifier. | |
| model_name: Embedding model name. | |
| texts: List of text strings (N). | |
| embeddings: (N, D) array of embedding vectors. | |
| metadatas: Optional per-entry metadata dicts. | |
| """ | |
| col_name = self._collection_name(session_id, model_name) | |
| collection = self._client.get_or_create_collection( | |
| name=col_name, | |
| metadata={"session_id": session_id, "model_name": model_name}, | |
| ) | |
| ids = [f"{session_id}_{i}" for i in range(len(texts))] | |
| if metadatas is None: | |
| metadatas = [{"index": i} for i in range(len(texts))] | |
| collection.upsert( | |
| ids=ids, | |
| documents=texts, | |
| embeddings=embeddings.tolist(), | |
| metadatas=metadatas, | |
| ) | |
| def load_embeddings( | |
| self, session_id: str, model_name: str | |
| ) -> Optional[np.ndarray]: | |
| """Load stored embeddings for a session+model pair. | |
| Returns (N, D) array or None if not found. | |
| """ | |
| col_name = self._collection_name(session_id, model_name) | |
| try: | |
| collection = self._client.get_collection(name=col_name) | |
| except Exception as e: | |
| logger.debug(f"Collection not found: {col_name}: {e}") | |
| return None | |
| result = collection.get(include=["embeddings"]) | |
| if result["embeddings"] is None or len(result["embeddings"]) == 0: | |
| return None | |
| return np.array(result["embeddings"], dtype=np.float32) | |
| def list_sessions(self) -> List[dict]: | |
| """List all stored session/model combinations.""" | |
| collections = self._client.list_collections() | |
| sessions = [] | |
| for col in collections: | |
| meta = col.metadata or {} | |
| sessions.append({ | |
| "collection_name": col.name, | |
| "session_id": meta.get("session_id", "unknown"), | |
| "model_name": meta.get("model_name", "unknown"), | |
| "count": col.count(), | |
| }) | |
| return sessions | |
| def delete_session(self, session_id: str, model_name: str): | |
| """Delete a stored session+model collection.""" | |
| col_name = self._collection_name(session_id, model_name) | |
| try: | |
| self._client.delete_collection(name=col_name) | |
| except Exception as e: | |
| logger.debug(f"Failed to delete collection {col_name}: {e}") | |
| def clear_all(self): | |
| """Delete all stored embedding collections.""" | |
| for col in self._client.list_collections(): | |
| try: | |
| self._client.delete_collection(name=col.name) | |
| except Exception as e: | |
| logger.debug(f"Failed to delete collection {col.name}: {e}") | |