fix: search; protect agaainst failed vector create/delete

This commit is contained in:
Rajat Ahuja 2025-12-05 15:24:07 -05:00
parent 1fb01a9ff8
commit 60fde82218
13 changed files with 594 additions and 69 deletions

View File

@ -1,10 +1,12 @@
"""add chunk_index to message_embeddings and make embeddings nullable
"""add chunk_index to message_embeddings, make embeddings nullable, add soft delete
This migration:
1. Adds the chunk_index column to message_embeddings table for tracking
chunked message embeddings in external vector stores (turbopuffer/lancedb).
2. Makes embedding columns nullable in both message_embeddings and documents tables
since embeddings are now stored in external vector stores instead of PostgreSQL.
3. Adds deleted_at column to documents table for soft delete support, enabling
hybrid sync/soft delete pattern for vector store consistency.
Revision ID: f1a2b3c4d5e6
Revises: baa22cad81e2
@ -30,7 +32,7 @@ schema = get_schema()
def upgrade() -> None:
"""Add chunk_index column to message_embeddings and make embeddings nullable."""
"""Add chunk_index, make embeddings nullable, add deleted_at for soft delete."""
inspector = sa.inspect(op.get_bind())
# Add chunk_index column to message_embeddings if it doesn't exist
@ -68,11 +70,40 @@ def upgrade() -> None:
schema=schema,
)
# Add deleted_at column to documents for soft delete support
# This enables hybrid sync/soft delete pattern:
# - Try to delete from vector store first
# - If successful, hard delete from DB
# - If vector delete fails, soft delete (set deleted_at) and let cleanup job handle it
if not column_exists("documents", "deleted_at", inspector):
op.add_column(
"documents",
sa.Column(
"deleted_at",
sa.DateTime(timezone=True),
nullable=True,
),
schema=schema,
)
# Create partial index for efficient cleanup queries (only index non-null values)
op.create_index(
"ix_documents_deleted_at",
"documents",
["deleted_at"],
schema=schema,
postgresql_where=sa.text("deleted_at IS NOT NULL"),
)
def downgrade() -> None:
"""Remove chunk_index column and revert embedding columns to non-nullable."""
"""Remove chunk_index, deleted_at columns and revert embedding columns."""
inspector = sa.inspect(op.get_bind())
# Remove deleted_at column and index from documents
if column_exists("documents", "deleted_at", inspector):
op.drop_index("ix_documents_deleted_at", table_name="documents", schema=schema)
op.drop_column("documents", "deleted_at", schema=schema)
# Revert documents.embedding back to nullable=True (it was originally nullable=True)
op.alter_column(
"documents",

View File

@ -1,11 +1,14 @@
import datetime
from collections.abc import Sequence
from logging import getLogger
from typing import Any
from sqlalchemy import delete, select
from sqlalchemy import delete, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql import Select
from sqlalchemy.sql.functions import func
from tenacity import AsyncRetrying, stop_after_attempt, wait_exponential
from src import models, schemas
from src.config import settings
@ -15,7 +18,7 @@ from src.crud.session import get_session
from src.embedding_client import embedding_client
from src.exceptions import ResourceNotFoundException, ValidationException
from src.utils.filter import apply_filter
from src.vector_store import VectorRecord, get_vector_store
from src.vector_store import VectorRecord, VectorStore, get_vector_store
logger = getLogger(__name__)
@ -50,6 +53,7 @@ def get_all_documents(
.where(models.Document.workspace_name == workspace_name)
.where(models.Document.observer == observer)
.where(models.Document.observed == observed)
.where(models.Document.deleted_at.is_(None)) # Exclude soft-deleted
)
# Apply additional filters if provided
@ -87,8 +91,10 @@ def get_documents_with_filters(
Returns:
Select query for documents
"""
stmt = select(models.Document).where(
models.Document.workspace_name == workspace_name
stmt = (
select(models.Document)
.where(models.Document.workspace_name == workspace_name)
.where(models.Document.deleted_at.is_(None)) # Exclude soft-deleted
)
# Apply additional filters if provided
@ -172,14 +178,17 @@ async def query_documents(
document_ids = [result.id for result in vector_results]
# Fetch documents from database
# No additional filtering needed since vector store already applied all supported filters
stmt = (
select(models.Document)
.where(models.Document.workspace_name == workspace_name)
.where(models.Document.observer == observer)
.where(models.Document.observed == observed)
.where(models.Document.deleted_at.is_(None))
.where(models.Document.id.in_(document_ids))
)
# Re-apply all filters at the database layer to catch any constraints
# that aren't supported by the vector store metadata.
stmt = apply_filter(stmt, models.Document, filters)
result = await db.execute(stmt)
documents = {doc.id: doc for doc in result.scalars().all()}
@ -281,7 +290,20 @@ async def create_documents(
},
)
)
await vector_store.upsert_many(namespace, vector_records)
# Retry vector upsert with exponential backoff
try:
async for attempt in AsyncRetrying(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
reraise=True,
):
with attempt:
await vector_store.upsert_many(namespace, vector_records)
except Exception as e:
# Final attempt failed - log but don't raise
# Documents exist in DB, vectors can be added manually later
logger.error(f"Failed to upsert vectors after retries: {e}")
except IntegrityError as e:
await db.rollback()
@ -302,7 +324,11 @@ async def delete_document(
session_name: str | None = None,
) -> None:
"""
Delete a single document by ID.
Delete a single document by ID using hybrid sync/soft delete pattern.
Tries to delete from vector store first, then hard deletes from DB.
If vector store delete fails, soft deletes (sets deleted_at) and lets
cleanup job handle vector deletion later.
Args:
db: Database session
@ -315,25 +341,53 @@ async def delete_document(
Raises:
ResourceNotFoundException: If document not found or doesn't match criteria
"""
stmt = delete(models.Document).where(
# Build base query conditions
conditions = [
models.Document.id == document_id,
models.Document.workspace_name == workspace_name,
models.Document.observer == observer,
models.Document.observed == observed,
)
# If session is specified, ensure document belongs to that session
models.Document.deleted_at.is_(None), # Only delete non-deleted docs
]
if session_name is not None:
stmt = stmt.where(models.Document.session_name == session_name)
conditions.append(models.Document.session_name == session_name)
result = await db.execute(stmt)
await db.commit()
# Check document exists first
check_stmt = select(models.Document).where(*conditions)
result = await db.execute(check_stmt)
doc = result.scalar_one_or_none()
if result.rowcount == 0:
if doc is None:
raise ResourceNotFoundException(
f"Document {document_id} not found or does not belong to the specified collection/session"
)
# Try to delete from vector store first
vector_store = get_vector_store()
namespace = vector_store.get_document_namespace(workspace_name, observer, observed)
vector_deleted = False
try:
await vector_store.delete_many(namespace, [document_id])
vector_deleted = True
except Exception as e:
logger.warning(f"Failed to delete vector for document {document_id}: {e}")
if vector_deleted:
# Happy path: hard delete from DB
delete_stmt = delete(models.Document).where(models.Document.id == document_id)
await db.execute(delete_stmt)
else:
# Fallback: soft delete, let cleanup job handle vector
update_stmt = (
update(models.Document)
.where(models.Document.id == document_id)
.values(deleted_at=func.now())
)
await db.execute(update_stmt)
await db.commit()
async def delete_document_by_id(
db: AsyncSession,
@ -341,7 +395,11 @@ async def delete_document_by_id(
document_id: str,
) -> None:
"""
Delete a single document by ID and workspace.
Delete a single document by ID and workspace using hybrid sync/soft delete pattern.
Tries to delete from vector store first, then hard deletes from DB.
If vector store delete fails, soft deletes (sets deleted_at) and lets
cleanup job handle vector deletion later.
Args:
db: Database session
@ -351,19 +409,48 @@ async def delete_document_by_id(
Raises:
ResourceNotFoundException: If document not found or doesn't belong to the workspace
"""
stmt = delete(models.Document).where(
# Fetch document to get observer/observed for namespace
stmt = select(models.Document).where(
models.Document.id == document_id,
models.Document.workspace_name == workspace_name,
models.Document.deleted_at.is_(None), # Only delete non-deleted docs
)
result = await db.execute(stmt)
await db.commit()
doc = result.scalar_one_or_none()
if result.rowcount == 0:
if doc is None:
raise ResourceNotFoundException(
f"Document {document_id} not found or does not belong to workspace {workspace_name}"
)
# Try to delete from vector store first
vector_store = get_vector_store()
namespace = vector_store.get_document_namespace(
workspace_name, doc.observer, doc.observed
)
vector_deleted = False
try:
await vector_store.delete_many(namespace, [document_id])
vector_deleted = True
except Exception as e:
logger.warning(f"Failed to delete vector for document {document_id}: {e}")
if vector_deleted:
# Happy path: hard delete from DB
delete_stmt = delete(models.Document).where(models.Document.id == document_id)
await db.execute(delete_stmt)
else:
# Fallback: soft delete, let cleanup job handle vector
update_stmt = (
update(models.Document)
.where(models.Document.id == document_id)
.values(deleted_at=func.now())
)
await db.execute(update_stmt)
await db.commit()
async def create_observations(
db: AsyncSession,
@ -479,7 +566,22 @@ async def create_observations(
},
)
)
await vector_store.upsert_many(namespace, vector_records)
# Retry vector upsert with exponential backoff
try:
async for attempt in AsyncRetrying(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
reraise=True,
):
with attempt:
await vector_store.upsert_many(namespace, vector_records)
except Exception as e:
# Final attempt failed - log but don't raise
# Documents exist in DB, vectors can be added manually later
logger.error(
f"Failed to upsert vectors for {namespace} after retries: {e}"
)
except IntegrityError as e:
await db.rollback()
@ -566,3 +668,89 @@ async def is_rejected_duplicate(
f"[DUPLICATE DETECTION] Rejecting new in favor of existing. new='{doc.content}', existing='{existing_doc.content}'."
)
return True
async def cleanup_soft_deleted_documents(
db: AsyncSession,
vector_store: VectorStore,
batch_size: int = 100,
older_than_minutes: int = 5,
) -> int:
"""
Clean up soft-deleted documents by deleting from vector store and hard deleting from DB.
This function is designed to be called periodically (e.g., every 5 minutes) to reconcile
any documents that were soft-deleted when the vector store was unavailable.
Steps:
1. Find documents with deleted_at older than threshold
2. Group by namespace (workspace/observer/observed)
3. Delete from vector store (per namespace)
4. Hard delete from DB only for documents where vector deletion succeeded
If vector deletion fails for a namespace, those documents remain soft-deleted
and will be retried on the next cleanup run.
Args:
db: Database session
vector_store: Vector store instance
batch_size: Maximum number of documents to process per call
older_than_minutes: Only process documents soft-deleted more than this many minutes ago
Returns:
Count of documents cleaned up (only those where vector deletion succeeded)
"""
cutoff = datetime.datetime.now(datetime.UTC) - datetime.timedelta(
minutes=older_than_minutes
)
# Find soft-deleted documents ready for cleanup
# Use FOR UPDATE SKIP LOCKED to prevent multiple deriver instances from
# processing the same documents simultaneously
stmt = (
select(models.Document)
.where(models.Document.deleted_at.is_not(None))
.where(models.Document.deleted_at < cutoff)
.limit(batch_size)
.with_for_update(skip_locked=True)
)
result = await db.execute(stmt)
documents = list(result.scalars().all())
if not documents:
return 0
# Group by namespace for batch vector deletion
by_namespace: dict[str, list[str]] = {}
for doc in documents:
namespace = vector_store.get_document_namespace(
doc.workspace_name, doc.observer, doc.observed
)
by_namespace.setdefault(namespace, []).append(doc.id)
# Delete from vector store (per namespace) and track successful deletions
successfully_deleted_ids: set[str] = set()
for namespace, ids in by_namespace.items():
try:
await vector_store.delete_many(namespace, ids)
# Only add to successfully_deleted_ids if vector deletion succeeded
successfully_deleted_ids.update(ids)
except Exception as e:
# Log but continue - vectors may already be deleted or namespace may not exist
logger.warning(f"Failed to delete vectors from {namespace}: {e}")
# Only hard delete documents where vector deletion succeeded
if successfully_deleted_ids:
await db.execute(
delete(models.Document).where(
models.Document.id.in_(successfully_deleted_ids)
)
)
await db.commit()
logger.debug(
f"Cleaned up {len(successfully_deleted_ids)} soft-deleted documents"
)
return len(successfully_deleted_ids)
# No documents were successfully deleted from vector store
return 0

View File

@ -4,6 +4,7 @@ from typing import Any
from nanoid import generate as generate_nanoid
from sqlalchemy import ColumnElement, Select, and_, func, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from tenacity import AsyncRetrying, stop_after_attempt, wait_exponential
from src import models, schemas
from src.config import settings
@ -179,9 +180,20 @@ async def create_messages(
db.add_all(embedding_objects)
await db.commit()
# Upsert vectors to external vector store
# Upsert vectors to external vector store with retry
if vector_records:
await vector_store.upsert_many(namespace, vector_records)
try:
async for attempt in AsyncRetrying(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
reraise=True,
):
with attempt:
await vector_store.upsert_many(namespace, vector_records)
except Exception as e:
# Final attempt failed - log but don't raise
# MessageEmbedding records exist in DB, vectors can be added later
logger.error(f"Failed to upsert message vectors after retries: {e}")
except Exception:
logger.exception(

View File

@ -18,6 +18,7 @@ from src.exceptions import (
ResourceNotFoundException,
)
from src.utils.filter import apply_filter
from src.vector_store import get_vector_store
from .peer import get_or_create_peers, get_peer
from .workspace import get_or_create_workspace
@ -420,6 +421,36 @@ async def delete_session(
)
)
# Delete message vectors from vector store before deleting DB records
# Fetch all MessageEmbedding records to build vector IDs
embedding_result = await db.execute(
select(
models.MessageEmbedding.message_id, models.MessageEmbedding.chunk_index
).where(
models.MessageEmbedding.session_name == session_name,
models.MessageEmbedding.workspace_name == workspace_name,
)
)
embeddings = embedding_result.all()
if embeddings:
# Build vector IDs: {message_id}_{chunk_index}
vector_ids = [f"{e.message_id}_{e.chunk_index}" for e in embeddings]
# Try to delete from vector store (best effort)
try:
vector_store = get_vector_store()
namespace = vector_store.get_message_namespace(workspace_name)
await vector_store.delete_many(namespace, vector_ids)
logger.debug(
f"Deleted {len(vector_ids)} message vectors for session {session_name}"
)
except Exception as e:
# Log warning but continue - workspace deletion will clean up eventually
logger.warning(
f"Failed to delete message vectors for session {session_name}: {e}"
)
# Delete MessageEmbedding entries in batches
await _batch_delete_matching(
db,
@ -431,6 +462,43 @@ async def delete_session(
batch_size=5000,
)
# Delete document vectors from vector store before deleting DB records
# Fetch all Document records to get IDs and namespaces
doc_result = await db.execute(
select(
models.Document.id,
models.Document.observer,
models.Document.observed,
).where(
models.Document.session_name == session_name,
models.Document.workspace_name == workspace_name,
)
)
documents = doc_result.all()
if documents:
# Group document IDs by namespace (observer/observed)
docs_by_namespace: dict[str, list[str]] = {}
vector_store = get_vector_store()
for doc in documents:
namespace = vector_store.get_document_namespace(
workspace_name, doc.observer, doc.observed
)
docs_by_namespace.setdefault(namespace, []).append(doc.id)
# Try to delete from vector store (best effort, per namespace)
for namespace, doc_ids in docs_by_namespace.items():
try:
await vector_store.delete_many(namespace, doc_ids)
logger.debug(
f"Deleted {len(doc_ids)} document vectors from {namespace}"
)
except Exception as e:
# Log warning but continue - workspace deletion will clean up eventually
logger.warning(
f"Failed to delete document vectors from {namespace}: {e}"
)
# Delete Document entries associated with this session in batches
await _batch_delete_matching(
db,

View File

@ -49,6 +49,10 @@ class WorkerOwnership(NamedTuple):
aqs_id: str # The ID of the ActiveQueueSession that the worker is processing
VECTOR_CLEANUP_INTERVAL_SECONDS = 300 # 5 minutes
QUEUE_CLEANUP_INTERVAL_SECONDS = 43200 # 12 hours
class QueueManager:
def __init__(self):
self.shutdown_event: asyncio.Event = asyncio.Event()
@ -111,7 +115,7 @@ class QueueManager:
)
logger.debug("Signal handlers registered")
# Start background maintenance loop
# Start background maintenance loop (handles both queue cleanup and vector cleanup)
try:
self._maintenance_task = asyncio.create_task(self._maintenance_loop())
except Exception:
@ -343,30 +347,85 @@ class QueueManager:
await db.commit()
async def _maintenance_loop(self) -> None:
"""Run periodic maintenance tasks on the queue."""
"""
Run periodic maintenance tasks.
- Vector cleanup: every 5 minutes (clean up soft-deleted documents)
- Queue cleanup: every 12 hours (remove old processed/errored queue items)
"""
# Track when each task should next run
next_vector_cleanup = datetime.now(timezone.utc)
next_queue_cleanup = datetime.now(timezone.utc)
try:
while not self.shutdown_event.is_set():
try:
await self.cleanup_queue_items()
except Exception:
logger.exception("Error during maintenance cleanup")
if settings.SENTRY.ENABLED:
sentry_sdk.capture_exception()
now = datetime.now(timezone.utc)
# Run vector cleanup if due
if now >= next_vector_cleanup:
try:
await self._run_vector_cleanup()
except Exception:
logger.exception("Error during vector cleanup")
if settings.SENTRY.ENABLED:
sentry_sdk.capture_exception()
next_vector_cleanup = now + timedelta(
seconds=VECTOR_CLEANUP_INTERVAL_SECONDS
)
# Run queue cleanup if due
if now >= next_queue_cleanup:
try:
await self.cleanup_queue_items()
except Exception:
logger.exception("Error during queue cleanup")
if settings.SENTRY.ENABLED:
sentry_sdk.capture_exception()
next_queue_cleanup = now + timedelta(
seconds=QUEUE_CLEANUP_INTERVAL_SECONDS
)
# Sleep until next task is due or shutdown
next_task_time = min(next_vector_cleanup, next_queue_cleanup)
sleep_seconds = max(
0, (next_task_time - datetime.now(timezone.utc)).total_seconds()
)
# Sleep until interval elapses or shutdown event is set
try:
await asyncio.wait_for(
self.shutdown_event.wait(),
timeout=43200, # 12 hours
timeout=sleep_seconds
or 1, # At least 1 second to avoid busy loop
)
break # Shutdown event set
except asyncio.TimeoutError:
# Timeout means it's time for next cleanup
# Timeout means it's time for next task
pass
except asyncio.CancelledError:
logger.debug("Maintenance loop cancelled")
raise
async def _run_vector_cleanup(self) -> None:
"""Run vector store cleanup for soft-deleted documents."""
from src.crud.document import cleanup_soft_deleted_documents
from src.vector_store import get_vector_store
async with tracked_db("vector_cleanup") as db:
vector_store = get_vector_store()
total_cleaned = 0
# Process in batches until no more soft-deleted documents
while True:
cleaned = await cleanup_soft_deleted_documents(db, vector_store)
total_cleaned += cleaned
if cleaned == 0:
break
if total_cleaned > 0:
logger.info(
f"Vector cleanup: removed {total_cleaned} soft-deleted documents"
)
async def _handle_processing_error(
self,
error: Exception,

View File

@ -389,6 +389,9 @@ class Document(Base):
ForeignKey("workspaces.name"), nullable=False, index=True
)
session_name: Mapped[str] = mapped_column(TEXT, index=True)
deleted_at: Mapped[datetime.datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True, index=True, default=None
)
collection = relationship("Collection", back_populates="documents")
__table_args__ = (

View File

@ -130,11 +130,11 @@ async def _semantic_search(
message_ids = list(seen_message_ids.keys())
# Fetch messages from database by the IDs from vector search
# No additional filtering needed since vector store already applied all filters
# Fetch messages from database by the IDs from vector search and reapply filters
semantic_query = select(models.Message).where(
models.Message.public_id.in_(message_ids)
)
semantic_query = apply_filter(semantic_query, models.Message, filters)
result = await db.execute(semantic_query)
messages = {msg.public_id: msg for msg in result.scalars().all()}

View File

@ -89,11 +89,9 @@ class VectorStore(ABC):
Args:
namespace: The namespace to store the vector in
id: Unique identifier for the vector
embedding: The embedding vector
metadata: Optional metadata to store with the vector
vector: VectorRecord containing id, embedding, and optional metadata
"""
pass
...
@abstractmethod
async def upsert_many(
@ -108,7 +106,7 @@ class VectorStore(ABC):
namespace: The namespace to store the vectors in
vectors: List of VectorRecord objects to upsert
"""
pass
...
@abstractmethod
async def query(
@ -133,7 +131,7 @@ class VectorStore(ABC):
Returns:
List of QueryResult objects, ordered by similarity (most similar first)
"""
pass
...
@abstractmethod
async def delete_many(self, namespace: str, ids: list[str]) -> None:
@ -144,7 +142,7 @@ class VectorStore(ABC):
namespace: The namespace containing the vectors
ids: List of vector identifiers to delete
"""
pass
...
@abstractmethod
async def delete_namespace(self, namespace: str) -> None:
@ -154,7 +152,7 @@ class VectorStore(ABC):
Args:
namespace: The namespace to delete
"""
pass
...
# Singleton instance

View File

@ -23,6 +23,8 @@ logger = logging.getLogger(__name__)
# Additional metadata columns are added dynamically
VECTOR_DIMENSION = 1536
# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false, reportUnknownParameterType=false
class LanceDBVectorStore(VectorStore):
"""
@ -35,7 +37,7 @@ class LanceDBVectorStore(VectorStore):
_db: AsyncConnection | None = None
_db_path: str
def __init__(self):
def __init__(self) -> None:
"""Initialize the LanceDB vector store."""
super().__init__()
self._db_path = settings.VECTOR_STORE.LANCEDB_PATH
@ -56,14 +58,14 @@ class LanceDBVectorStore(VectorStore):
return None
async def _get_or_create_table(
self, namespace: str, sample_data: list[dict[str, Any]] | None = None
self,
namespace: str,
) -> AsyncTable:
"""
Get existing table or create if not exists.
Args:
namespace: Table name (namespace)
sample_data: Optional sample data to infer schema from
Returns:
LanceDB async table
@ -73,18 +75,46 @@ class LanceDBVectorStore(VectorStore):
if namespace in table_names:
return await db.open_table(namespace)
# Create table with sample data if provided
if sample_data:
return await db.create_table(namespace, data=sample_data)
# Create empty table with base schema
schema = pa.schema( # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType]
[
pa.field("id", pa.string()), # pyright: ignore[reportUnknownMemberType]
pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)), # pyright: ignore[reportUnknownMemberType]
fields: list[pa.Field] = [
pa.field("id", pa.string()),
pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)),
]
fields.extend(self._metadata_fields_for_namespace(namespace))
schema = pa.schema(fields)
table = await db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType]
return table
def _metadata_fields_for_namespace(self, namespace: str) -> list[pa.Field]:
"""
Infer standard metadata columns based on namespace structure.
Namespaces:
- Documents: {prefix}.{workspace}.{observer}.{observed}
- Messages: {prefix}.{workspace}.messages
"""
parts = namespace.split(".")
if len(parts) < 3:
return []
if parts[-1] == "messages":
return [
pa.field("message_id", pa.string(), nullable=True),
pa.field("session_name", pa.string(), nullable=True),
pa.field("peer_name", pa.string(), nullable=True),
pa.field("chunk_index", pa.int64(), nullable=True),
]
)
return await db.create_table(namespace, schema=schema) # pyright: ignore[reportUnknownArgumentType]
if len(parts) == 4:
return [
pa.field("workspace_name", pa.string(), nullable=True),
pa.field("observer", pa.string(), nullable=True),
pa.field("observed", pa.string(), nullable=True),
pa.field("session_name", pa.string(), nullable=True),
pa.field("level", pa.string(), nullable=True),
]
return []
def _row_to_dict(self, vector: VectorRecord) -> dict[str, Any]:
"""Convert a VectorRecord to a dict for LanceDB."""
@ -94,7 +124,10 @@ class LanceDBVectorStore(VectorStore):
}
# Add metadata fields
if vector.metadata:
row.update(vector.metadata)
reserved_keys = {"id", "vector", "_distance"}
for key in vector.metadata:
if key not in reserved_keys:
row[key] = vector.metadata[key]
return row
async def upsert(
@ -111,7 +144,7 @@ class LanceDBVectorStore(VectorStore):
"""
try:
row = self._row_to_dict(vector)
table = await self._get_or_create_table(namespace, sample_data=[row])
table = await self._get_or_create_table(namespace)
# Use merge_insert for upsert behavior
await (
@ -145,7 +178,7 @@ class LanceDBVectorStore(VectorStore):
try:
rows = [self._row_to_dict(v) for v in vectors]
table = await self._get_or_create_table(namespace, sample_data=rows)
table = await self._get_or_create_table(namespace)
# Use merge_insert for upsert behavior
await (

View File

@ -70,9 +70,7 @@ class TurbopufferVectorStore(VectorStore):
Args:
namespace: The namespace to store the vector in
id: Unique identifier for the vector
embedding: The embedding vector
metadata: Optional metadata to store with the vector
vector: VectorRecord containing id, embedding, and optional metadata
"""
ns = self._get_namespace(namespace)
attributes = vector.metadata or {}
@ -116,7 +114,7 @@ class TurbopufferVectorStore(VectorStore):
{
"id": v.id,
"vector": v.embedding,
**v.metadata,
**(v.metadata or {}),
}
for v in vectors
]

View File

@ -22,7 +22,7 @@ from . import (
test_d429de0e5338_adopt_peer_paradigm,
test_e9b705f9adf9_add_server_defaults_to_timestamp_,
test_ec8f94139b02_codify_workspace_name_and_message_id_in_,
test_f1a2b3c4d5e6_make_embeddings_nullable,
test_f1a2b3c4d5e6_support_external_embeddings,
)
__all__ = [
@ -47,5 +47,5 @@ __all__ = [
"test_d429de0e5338_adopt_peer_paradigm",
"test_e9b705f9adf9_add_server_defaults_to_timestamp_",
"test_ec8f94139b02_codify_workspace_name_and_message_id_in_",
"test_f1a2b3c4d5e6_make_embeddings_nullable",
"test_f1a2b3c4d5e6_support_external_embeddings",
]

View File

@ -1,3 +1,5 @@
import datetime
import pytest
from nanoid import generate as generate_nanoid
from sqlalchemy import select
@ -139,6 +141,139 @@ class TestDocumentCRUD:
assert len(results) == 2
@pytest.mark.asyncio
async def test_query_documents_excludes_soft_deleted(
self,
db_session: AsyncSession,
sample_data: tuple[models.Workspace, models.Peer],
):
"""Query results should not include soft-deleted documents even if vectors remain"""
test_workspace, test_peer = sample_data
test_peer2, test_session, _ = await self._setup_test_data(
db_session, test_workspace, test_peer
)
# Create two documents and persist embeddings
doc_schemas = [
schemas.DocumentCreate(
content="User likes pizza",
embedding=[0.9] * 1536,
session_name=test_session.name,
metadata=schemas.DocumentMetadata(
message_ids=[1],
message_created_at="2025-01-01T00:00:00Z",
),
),
schemas.DocumentCreate(
content="User dislikes vegetables",
embedding=[0.1] * 1536,
session_name=test_session.name,
metadata=schemas.DocumentMetadata(
message_ids=[2],
message_created_at="2025-01-01T00:00:00Z",
),
),
]
await crud.create_documents(
db_session,
doc_schemas,
workspace_name=test_workspace.name,
observer=test_peer.name,
observed=test_peer2.name,
)
# Soft-delete one document without touching vectors
stmt = select(models.Document).where(
models.Document.workspace_name == test_workspace.name,
models.Document.observer == test_peer.name,
models.Document.observed == test_peer2.name,
)
result = await db_session.execute(stmt)
docs = {doc.content: doc for doc in result.scalars().all()}
deleted_doc = docs["User likes pizza"]
kept_doc = docs["User dislikes vegetables"]
deleted_doc.deleted_at = datetime.datetime.now(datetime.timezone.utc)
await db_session.commit()
results = await crud.query_documents(
db_session,
workspace_name=test_workspace.name,
query="food preferences",
observer=test_peer.name,
observed=test_peer2.name,
top_k=10,
)
assert len(results) == 1
assert results[0].id == kept_doc.id
@pytest.mark.asyncio
async def test_query_documents_applies_additional_filters(
self,
db_session: AsyncSession,
sample_data: tuple[models.Workspace, models.Peer],
):
"""Filters beyond vector metadata should be enforced at the DB layer"""
test_workspace, test_peer = sample_data
test_peer2, test_session, _ = await self._setup_test_data(
db_session, test_workspace, test_peer
)
doc_schemas = [
schemas.DocumentCreate(
content="Observation one",
embedding=[0.5] * 1536,
session_name=test_session.name,
times_derived=1,
metadata=schemas.DocumentMetadata(
message_ids=[1],
message_created_at="2025-01-01T00:00:00Z",
),
),
schemas.DocumentCreate(
content="Observation two",
embedding=[0.5] * 1536,
session_name=test_session.name,
times_derived=2,
metadata=schemas.DocumentMetadata(
message_ids=[2],
message_created_at="2025-01-01T00:00:00Z",
),
),
]
await crud.create_documents(
db_session,
doc_schemas,
workspace_name=test_workspace.name,
observer=test_peer.name,
observed=test_peer2.name,
)
result = await db_session.execute(
select(models.Document).where(
models.Document.workspace_name == test_workspace.name,
models.Document.observer == test_peer.name,
models.Document.observed == test_peer2.name,
)
)
docs = result.scalars().all()
times_derived_map = {doc.times_derived: doc.id for doc in docs}
results = await crud.query_documents(
db_session,
workspace_name=test_workspace.name,
query="any query",
observer=test_peer.name,
observed=test_peer2.name,
top_k=10,
filters={"times_derived": 2},
embedding=[0.5] * 1536,
)
assert len(results) == 1
assert results[0].id == times_derived_map[2]
@pytest.mark.asyncio
async def test_delete_document_success(
self,