diff --git a/src/crud/document.py b/src/crud/document.py index ae058cf9..7bcb3497 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -293,7 +293,7 @@ async def is_rejected_duplicate( namespace = vector_store.get_document_namespace( workspace_name, observer, observed ) - await vector_store.delete(namespace, existing_doc.id) + await vector_store.delete_many(namespace, [existing_doc.id]) return False # Don't reject the new document diff --git a/src/utils/search.py b/src/utils/search.py index 268c6b81..2f3bc4f1 100644 --- a/src/utils/search.py +++ b/src/utils/search.py @@ -238,8 +238,9 @@ async def search( # Perform semantic search if enabled and we have workspace context # workspace_id is required for semantic search to determine the vector namespace - workspace_name = filters.get("workspace_id") if filters else None - if settings.EMBED_MESSAGES and workspace_name: + workspace_name: str | None = filters.get("workspace_id") if filters else None + if settings.EMBED_MESSAGES and isinstance(workspace_name, str): + # Type narrowing: workspace_name is guaranteed to be str in this block # Get more results for fusion semantic_limit = limit * 2 semantic_results = await _semantic_search( diff --git a/src/vector_store/__init__.py b/src/vector_store/__init__.py index ce6d9409..5437537a 100644 --- a/src/vector_store/__init__.py +++ b/src/vector_store/__init__.py @@ -7,8 +7,6 @@ from dataclasses import dataclass, field from typing import Any from src.config import settings -from src.vector_store.lancedb import LanceDBVectorStore -from src.vector_store.turbopuffer import TurbopufferVectorStore @dataclass @@ -34,8 +32,11 @@ class VectorStore(ABC): Abstract base class for vector store implementations. All vector operations are namespace-scoped. Namespaces map to: - - Document embeddings: {prefix}-{workspace}-{observer}-{observed} (per collection) - - Message embeddings: {prefix}-{workspace}-messages (per workspace) + - Document embeddings: {prefix}:{workspace}:{observer}:{observed} (per collection) + - Message embeddings: {prefix}:{workspace}:messages (per workspace) + + Note: Colon (:) is used as the delimiter to avoid collisions since it's not + allowed in workspace/peer IDs (which only allow [A-Za-z0-9_-]). """ namespace_prefix: str @@ -59,9 +60,9 @@ class VectorStore(ABC): observed: Name of the observed peer Returns: - Namespace string in format: {prefix}_{workspace}_{observer}_{observed} + Namespace string in format: {prefix}:{workspace}:{observer}:{observed} """ - return f"{self.namespace_prefix}_{workspace_name}_{observer}_{observed}" + return f"{self.namespace_prefix}:{workspace_name}:{observer}:{observed}" def get_message_namespace(self, workspace_name: str) -> str: """ @@ -71,9 +72,9 @@ class VectorStore(ABC): workspace_name: Name of the workspace Returns: - Namespace string in format: {prefix}_{workspace}_messages + Namespace string in format: {prefix}:{workspace}:messages """ - return f"{self.namespace_prefix}_{workspace_name}_messages" + return f"{self.namespace_prefix}:{workspace_name}:messages" # === Core operations === @abstractmethod @@ -177,8 +178,12 @@ def get_vector_store() -> VectorStore: store_type = settings.VECTOR_STORE.TYPE if store_type == "turbopuffer": + from src.vector_store.turbopuffer import TurbopufferVectorStore + _vector_store_instance = TurbopufferVectorStore() elif store_type == "lancedb": + from src.vector_store.lancedb import LanceDBVectorStore + _vector_store_instance = LanceDBVectorStore() else: raise ValueError(f"Unknown vector store type: {store_type}") diff --git a/src/vector_store/lancedb.py b/src/vector_store/lancedb.py index a5735d79..49f05c4f 100644 --- a/src/vector_store/lancedb.py +++ b/src/vector_store/lancedb.py @@ -65,13 +65,13 @@ class LanceDBVectorStore(VectorStore): return self._db.create_table(namespace, data=sample_data) # Create empty table with base schema - schema = pa.schema( + schema = pa.schema( # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] [ - pa.field("id", pa.string()), - pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)), + pa.field("id", pa.string()), # pyright: ignore[reportUnknownMemberType] + pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)), # pyright: ignore[reportUnknownMemberType] ] ) - return self._db.create_table(namespace, schema=schema) + return self._db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType] def _row_to_dict(self, vector: VectorRecord) -> dict[str, Any]: """Convert a VectorRecord to a dict for LanceDB.""" @@ -168,36 +168,37 @@ class LanceDBVectorStore(VectorStore): try: # Build query (LanceDB types are incomplete, so type checker reports false positive) - query = table.search(embedding).distance_type("cosine").limit(top_k) # pyright: ignore[reportAttributeAccessIssue] + query = table.search(embedding).distance_type("cosine").limit(top_k) # pyright: ignore[reportAttributeAccessIssue, reportUnknownVariableType, reportUnknownMemberType] # Apply filters if provided if filters: where_clause = self._build_where_clause(filters) if where_clause: - query = query.where(where_clause) + query = query.where(where_clause) # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # Execute query - results = query.to_list() + results = query.to_list() # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # Convert to QueryResult objects query_results: list[QueryResult] = [] - for row in results: - dist = float(row.get("_distance", 0.0)) + for row in results: # pyright: ignore[reportUnknownVariableType] + dist = float(row.get("_distance", 0.0)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Filter by max_distance if specified if max_distance is not None and dist > max_distance: continue # Extract metadata (everything except id, vector, _distance) + # Type annotations for dict comprehension to satisfy type checker metadata: dict[str, Any] = { k: v - for k, v in row.items() + for k, v in row.items() # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] if k not in ("id", "vector", "_distance") } query_results.append( QueryResult( - id=str(row["id"]), + id=str(row["id"]), # pyright: ignore[reportUnknownArgumentType] score=dist, metadata=metadata, )