fix: bug fixes
This commit is contained in:
parent
0dd3ee4e45
commit
2af9b1e1ca
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Reference in New Issue