fix: bug fixes

This commit is contained in:
Rajat Ahuja 2025-12-04 14:26:46 -05:00
parent 0dd3ee4e45
commit 2af9b1e1ca
4 changed files with 29 additions and 22 deletions

View File

@ -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

View File

@ -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(

View File

@ -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}")

View File

@ -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,
)