fix: set up compose vector store and reconciliation loop
This commit is contained in:
parent
1988f18fbe
commit
b0051c009e
|
|
@ -1,15 +1,15 @@
|
|||
"""add chunk_index to message_embeddings, make embeddings nullable, add soft delete
|
||||
"""make embeddings nullable, add soft delete, add vector sync state
|
||||
|
||||
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
|
||||
1. 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
|
||||
2. Adds deleted_at column to documents table for soft delete support, enabling
|
||||
hybrid sync/soft delete pattern for vector store consistency.
|
||||
3. Adds sync_state, last_sync_at, and sync_attempts columns to documents and
|
||||
message_embeddings tables for tracking vector store synchronization status.
|
||||
|
||||
Revision ID: f1a2b3c4d5e6
|
||||
Revises: baa22cad81e2
|
||||
Revises: 110bdf470272
|
||||
Create Date: 2025-11-24 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
|
@ -24,7 +24,7 @@ from migrations.utils import column_exists, get_schema
|
|||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f1a2b3c4d5e6"
|
||||
down_revision: str | None = "baa22cad81e2"
|
||||
down_revision: str | None = "110bdf470272"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
|
@ -32,26 +32,9 @@ schema = get_schema()
|
|||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add chunk_index, make embeddings nullable, add deleted_at for soft delete."""
|
||||
"""Make embeddings nullable, add deleted_at, add sync state."""
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
|
||||
# Add chunk_index column to message_embeddings if it doesn't exist
|
||||
# This is needed to track which chunk of a message this embedding represents
|
||||
# Vector ID format: {message_public_id}_{chunk_index}
|
||||
if not column_exists("message_embeddings", "chunk_index", inspector):
|
||||
op.add_column(
|
||||
"message_embeddings",
|
||||
sa.Column(
|
||||
"chunk_index",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Make message_embeddings.embedding nullable since embeddings are now stored
|
||||
# in external vector stores (turbopuffer/lancedb) instead of PostgreSQL
|
||||
op.alter_column(
|
||||
"message_embeddings",
|
||||
"embedding",
|
||||
|
|
@ -60,8 +43,6 @@ def upgrade() -> None:
|
|||
schema=schema,
|
||||
)
|
||||
|
||||
# Make documents.embedding nullable for the same reason
|
||||
# (this should already be nullable, but ensure it for consistency)
|
||||
op.alter_column(
|
||||
"documents",
|
||||
"embedding",
|
||||
|
|
@ -71,10 +52,6 @@ def upgrade() -> None:
|
|||
)
|
||||
|
||||
# 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",
|
||||
|
|
@ -94,35 +71,126 @@ def upgrade() -> None:
|
|||
postgresql_where=sa.text("deleted_at IS NOT NULL"),
|
||||
)
|
||||
|
||||
# Add sync state columns to documents table
|
||||
if not column_exists("documents", "sync_state", inspector):
|
||||
op.add_column(
|
||||
"documents",
|
||||
sa.Column(
|
||||
"sync_state",
|
||||
sa.TEXT(),
|
||||
nullable=False,
|
||||
server_default="pending", # Existing records need reconciliation
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_documents_sync_state",
|
||||
"documents",
|
||||
["sync_state"],
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if not column_exists("documents", "last_sync_at", inspector):
|
||||
op.add_column(
|
||||
"documents",
|
||||
sa.Column(
|
||||
"last_sync_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=True,
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if not column_exists("documents", "sync_attempts", inspector):
|
||||
op.add_column(
|
||||
"documents",
|
||||
sa.Column(
|
||||
"sync_attempts",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Add sync state columns to message_embeddings table
|
||||
if not column_exists("message_embeddings", "sync_state", inspector):
|
||||
op.add_column(
|
||||
"message_embeddings",
|
||||
sa.Column(
|
||||
"sync_state",
|
||||
sa.TEXT(),
|
||||
nullable=False,
|
||||
server_default="pending", # Existing records need reconciliation
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_message_embeddings_sync_state",
|
||||
"message_embeddings",
|
||||
["sync_state"],
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if not column_exists("message_embeddings", "last_sync_at", inspector):
|
||||
op.add_column(
|
||||
"message_embeddings",
|
||||
sa.Column(
|
||||
"last_sync_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=True,
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if not column_exists("message_embeddings", "sync_attempts", inspector):
|
||||
op.add_column(
|
||||
"message_embeddings",
|
||||
sa.Column(
|
||||
"sync_attempts",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove chunk_index, deleted_at columns and revert embedding columns."""
|
||||
"""Remove deleted_at columns and revert embedding columns."""
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
|
||||
# Remove sync state columns from message_embeddings
|
||||
if column_exists("message_embeddings", "sync_attempts", inspector):
|
||||
op.drop_column("message_embeddings", "sync_attempts", schema=schema)
|
||||
|
||||
if column_exists("message_embeddings", "last_sync_at", inspector):
|
||||
op.drop_column("message_embeddings", "last_sync_at", schema=schema)
|
||||
|
||||
if column_exists("message_embeddings", "sync_state", inspector):
|
||||
op.drop_index(
|
||||
"ix_message_embeddings_sync_state",
|
||||
table_name="message_embeddings",
|
||||
schema=schema,
|
||||
)
|
||||
op.drop_column("message_embeddings", "sync_state", schema=schema)
|
||||
|
||||
# Remove sync state columns from documents
|
||||
if column_exists("documents", "sync_attempts", inspector):
|
||||
op.drop_column("documents", "sync_attempts", schema=schema)
|
||||
|
||||
if column_exists("documents", "last_sync_at", inspector):
|
||||
op.drop_column("documents", "last_sync_at", schema=schema)
|
||||
|
||||
if column_exists("documents", "sync_state", inspector):
|
||||
op.drop_index("ix_documents_sync_state", table_name="documents", schema=schema)
|
||||
op.drop_column("documents", "sync_state", schema=schema)
|
||||
|
||||
# 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",
|
||||
"embedding",
|
||||
existing_type=Vector(1536),
|
||||
nullable=True, # Keep as nullable since it was nullable in the original schema
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Revert message_embeddings.embedding back to nullable=False
|
||||
# Note: This may fail if there are NULL values in the database
|
||||
op.alter_column(
|
||||
"message_embeddings",
|
||||
"embedding",
|
||||
existing_type=Vector(1536),
|
||||
nullable=False,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Remove chunk_index column if it exists
|
||||
if column_exists("message_embeddings", "chunk_index", inspector):
|
||||
op.drop_column("message_embeddings", "chunk_index", schema=schema)
|
||||
# NOTE: This downgrade does NOT restore the NOT NULL constraint on embedding columns
|
||||
# in message_embeddings and documents tables. This is intentional to avoid migration
|
||||
# failures if NULL embedding values exist (which is expected when using external vector stores).
|
||||
|
|
|
|||
|
|
@ -358,19 +358,33 @@ class DreamSettings(BackupLLMSettingsMixin, HonchoSettings):
|
|||
|
||||
|
||||
class VectorStoreSettings(HonchoSettings):
|
||||
"""Settings for external vector store (Turbopuffer or LanceDB)."""
|
||||
"""Settings for vector store (pgvector, Turbopuffer, or LanceDB)."""
|
||||
|
||||
model_config = SettingsConfigDict(env_prefix="VECTOR_STORE_", extra="ignore") # pyright: ignore
|
||||
|
||||
# Vector store type: "turbopuffer" or "lancedb"
|
||||
TYPE: Literal["turbopuffer", "lancedb"] = "lancedb"
|
||||
# Primary vector store type
|
||||
PRIMARY_TYPE: Literal["pgvector", "turbopuffer", "lancedb"] = "pgvector"
|
||||
|
||||
# Secondary vector store type (optional)
|
||||
# When set, enables:
|
||||
# - Dual-write: writes go to both primary and secondary
|
||||
# - Fallback read: reads try primary first, fall back to secondary if empty
|
||||
SECONDARY_TYPE: Literal["pgvector", "turbopuffer", "lancedb"] | None = None
|
||||
|
||||
# Global namespace prefix for all vector namespaces
|
||||
# Namespaces follow the pattern:
|
||||
# - Documents: {NAMESPACE}-{workspace}-{observer}-{observed}
|
||||
# - Messages: {NAMESPACE}-{workspace}-messages
|
||||
# - Documents: {NAMESPACE}.{workspace}.{observer}.{observed}
|
||||
# - Messages: {NAMESPACE}.{workspace}.messages
|
||||
NAMESPACE: str = "honcho"
|
||||
|
||||
DIMENSIONS: Annotated[
|
||||
int,
|
||||
Field(
|
||||
default=1536,
|
||||
gt=0,
|
||||
),
|
||||
] = 1536
|
||||
|
||||
# Turbopuffer-specific settings
|
||||
TURBOPUFFER_API_KEY: str | None = None
|
||||
TURBOPUFFER_REGION: str | None = None
|
||||
|
|
@ -380,12 +394,30 @@ class VectorStoreSettings(HonchoSettings):
|
|||
|
||||
@model_validator(mode="after")
|
||||
def _require_api_key_for_turbopuffer(self) -> "VectorStoreSettings":
|
||||
if self.TYPE == "turbopuffer" and not self.TURBOPUFFER_API_KEY:
|
||||
if self.PRIMARY_TYPE == "turbopuffer" and not self.TURBOPUFFER_API_KEY:
|
||||
raise ValueError(
|
||||
"VECTOR_STORE_TURBOPUFFER_API_KEY must be set when TYPE is 'turbopuffer'"
|
||||
"VECTOR_STORE_TURBOPUFFER_API_KEY must be set when PRIMARY_TYPE is 'turbopuffer'"
|
||||
)
|
||||
if self.SECONDARY_TYPE == "turbopuffer" and not self.TURBOPUFFER_API_KEY:
|
||||
raise ValueError(
|
||||
"VECTOR_STORE_TURBOPUFFER_API_KEY must be set when SECONDARY_TYPE is 'turbopuffer'"
|
||||
)
|
||||
return self
|
||||
|
||||
@property
|
||||
def should_run_reconciliation(self) -> bool:
|
||||
"""
|
||||
Determine if vector reconciliation should run.
|
||||
|
||||
Reconciliation syncs embeddings from postgres (used by pgvector) to
|
||||
external vector stores. It only runs when:
|
||||
1. A secondary store is configured AND
|
||||
2. pgvector is involved as either primary or secondary
|
||||
"""
|
||||
return self.SECONDARY_TYPE is not None and (
|
||||
self.PRIMARY_TYPE == "pgvector" or self.SECONDARY_TYPE == "pgvector"
|
||||
)
|
||||
|
||||
|
||||
class AppSettings(HonchoSettings):
|
||||
# No env_prefix for app-level settings
|
||||
|
|
@ -445,13 +477,11 @@ class AppSettings(HonchoSettings):
|
|||
VECTOR_STORE.NAMESPACE are guaranteed to exist. Explicitly provided
|
||||
nested namespaces are preserved.
|
||||
"""
|
||||
if self.CACHE.NAMESPACE is None:
|
||||
if "NAMESPACE" not in self.CACHE.model_fields_set:
|
||||
self.CACHE.NAMESPACE = self.NAMESPACE
|
||||
if self.METRICS.NAMESPACE is None:
|
||||
if "NAMESPACE" not in self.METRICS.model_fields_set:
|
||||
self.METRICS.NAMESPACE = self.NAMESPACE
|
||||
|
||||
vector_namespace_explicit = "NAMESPACE" in self.VECTOR_STORE.model_fields_set
|
||||
if not vector_namespace_explicit:
|
||||
if "NAMESPACE" not in self.VECTOR_STORE.model_fields_set:
|
||||
self.VECTOR_STORE.NAMESPACE = self.NAMESPACE
|
||||
|
||||
return self
|
||||
|
|
|
|||
|
|
@ -8,7 +8,13 @@ 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 tenacity import (
|
||||
AsyncRetrying,
|
||||
retry_if_exception_type,
|
||||
retry_if_result,
|
||||
stop_after_attempt,
|
||||
wait_exponential,
|
||||
)
|
||||
|
||||
from src import models, schemas
|
||||
from src.config import settings
|
||||
|
|
@ -149,7 +155,9 @@ async def query_documents(
|
|||
|
||||
# Get vector store and namespace for this collection
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(workspace_name, observer, observed)
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document", workspace_name, observer, observed
|
||||
)
|
||||
|
||||
# Build vector store filters
|
||||
# Convert filter dict to vector store format (handles level, session_name, etc.)
|
||||
|
|
@ -241,16 +249,42 @@ async def create_documents(
|
|||
continue
|
||||
|
||||
metadata_dict = doc.metadata.model_dump(exclude_none=True)
|
||||
new_doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=observer,
|
||||
observed=observed,
|
||||
content=doc.content,
|
||||
level=doc.level,
|
||||
times_derived=doc.times_derived,
|
||||
internal_metadata=metadata_dict,
|
||||
session_name=doc.session_name,
|
||||
|
||||
# Check if pgvector is being used (primary or secondary)
|
||||
# If so, write embeddings to ORM since pgvector relies on postgres
|
||||
pgvector_in_use = (
|
||||
settings.VECTOR_STORE.PRIMARY_TYPE == "pgvector"
|
||||
or settings.VECTOR_STORE.SECONDARY_TYPE == "pgvector"
|
||||
)
|
||||
|
||||
if pgvector_in_use and doc.embedding:
|
||||
# pgvector in use: write embedding to ORM (postgres)
|
||||
new_doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=observer,
|
||||
observed=observed,
|
||||
content=doc.content,
|
||||
level=doc.level,
|
||||
times_derived=doc.times_derived,
|
||||
internal_metadata=metadata_dict,
|
||||
session_name=doc.session_name,
|
||||
embedding=doc.embedding,
|
||||
)
|
||||
else:
|
||||
# pgvector not in use or no embedding: don't write embedding to postgres
|
||||
new_doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=observer,
|
||||
observed=observed,
|
||||
content=doc.content,
|
||||
level=doc.level,
|
||||
times_derived=doc.times_derived,
|
||||
internal_metadata=metadata_dict,
|
||||
session_name=doc.session_name,
|
||||
)
|
||||
|
||||
if doc.embedding:
|
||||
new_doc.sync_state = "pending"
|
||||
honcho_documents.append(new_doc)
|
||||
|
||||
# Track embedding for vector store (ID will be available after commit)
|
||||
|
|
@ -270,13 +304,18 @@ async def create_documents(
|
|||
# Store embeddings in vector store after documents are committed (IDs now available)
|
||||
if docs_with_embeddings:
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, observer, observed
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
observer,
|
||||
observed,
|
||||
)
|
||||
|
||||
# Build vector records with metadata for filtering
|
||||
vector_records: list[VectorRecord] = []
|
||||
doc_ids: list[str] = []
|
||||
for doc, embedding in docs_with_embeddings:
|
||||
doc_ids.append(doc.id)
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=doc.id,
|
||||
|
|
@ -291,19 +330,64 @@ async def create_documents(
|
|||
)
|
||||
)
|
||||
|
||||
# Retry vector upsert with exponential backoff
|
||||
# Retry vector upsert with exponential backoff (3 attempts)
|
||||
try:
|
||||
result = None
|
||||
async for attempt in AsyncRetrying(
|
||||
stop=stop_after_attempt(3),
|
||||
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
|
||||
retry=retry_if_exception_type(Exception)
|
||||
| retry_if_result(
|
||||
lambda res: res is not None and res.secondary_ok is False
|
||||
),
|
||||
reraise=True,
|
||||
):
|
||||
with attempt:
|
||||
await vector_store.upsert_many(namespace, vector_records)
|
||||
result = await vector_store.upsert_many(
|
||||
namespace, vector_records
|
||||
)
|
||||
|
||||
if result is not None and result.secondary_ok is False:
|
||||
# Partial success: primary has data but secondary doesn't
|
||||
# Keep as "pending" for reconciliation to sync secondary
|
||||
logger.warning(
|
||||
f"Partial sync for namespace {namespace}: {result.secondary_error}"
|
||||
)
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_attempts=models.Document.sync_attempts + 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
else:
|
||||
# Success: both primary and secondary stores have the data
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_state="synced",
|
||||
last_sync_at=func.now(),
|
||||
sync_attempts=0,
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
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}")
|
||||
# Total failure: primary write failed
|
||||
# Keep as "pending" for reconciliation to retry
|
||||
logger.error(f"Failed to upsert vectors after 3 retries: {e}")
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_attempts=models.Document.sync_attempts + 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
except IntegrityError as e:
|
||||
await db.rollback()
|
||||
|
|
@ -364,7 +448,9 @@ async def delete_document(
|
|||
|
||||
# Try to delete from vector store first
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(workspace_name, observer, observed)
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document", workspace_name, observer, observed
|
||||
)
|
||||
vector_deleted = False
|
||||
|
||||
try:
|
||||
|
|
@ -425,8 +511,11 @@ async def delete_document_by_id(
|
|||
|
||||
# Try to delete from vector store first
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, doc.observer, doc.observed
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
doc.observer,
|
||||
doc.observed,
|
||||
)
|
||||
vector_deleted = False
|
||||
|
||||
|
|
@ -517,17 +606,40 @@ async def create_observations(
|
|||
tuple[str, str], list[tuple[models.Document, list[float]]]
|
||||
] = {}
|
||||
|
||||
# Check if pgvector is being used (primary or secondary)
|
||||
# If so, write embeddings to ORM since pgvector relies on postgres
|
||||
pgvector_in_use = (
|
||||
settings.VECTOR_STORE.PRIMARY_TYPE == "pgvector"
|
||||
or settings.VECTOR_STORE.SECONDARY_TYPE == "pgvector"
|
||||
)
|
||||
|
||||
for obs, embedding in zip(observations, embeddings, strict=True):
|
||||
doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=obs.observer_id,
|
||||
observed=obs.observed_id,
|
||||
content=obs.content,
|
||||
level="explicit", # Manually created observations are always explicit
|
||||
times_derived=1,
|
||||
internal_metadata={}, # No message_ids since not derived from messages
|
||||
session_name=obs.session_id,
|
||||
)
|
||||
if pgvector_in_use:
|
||||
# pgvector in use: write embedding to ORM (postgres)
|
||||
doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=obs.observer_id,
|
||||
observed=obs.observed_id,
|
||||
content=obs.content,
|
||||
level="explicit", # Manually created observations are always explicit
|
||||
times_derived=1,
|
||||
internal_metadata={}, # No message_ids since not derived from messages
|
||||
session_name=obs.session_id,
|
||||
embedding=embedding,
|
||||
)
|
||||
else:
|
||||
# pgvector not in use: don't write embedding to postgres
|
||||
doc = models.Document(
|
||||
workspace_name=workspace_name,
|
||||
observer=obs.observer_id,
|
||||
observed=obs.observed_id,
|
||||
content=obs.content,
|
||||
level="explicit", # Manually created observations are always explicit
|
||||
times_derived=1,
|
||||
internal_metadata={}, # No message_ids since not derived from messages
|
||||
session_name=obs.session_id,
|
||||
)
|
||||
doc.sync_state = "pending"
|
||||
honcho_documents.append(doc)
|
||||
|
||||
# Track embedding for vector store (grouped by collection)
|
||||
|
|
@ -546,13 +658,18 @@ async def create_observations(
|
|||
# Store embeddings in vector store after documents are committed (IDs now available)
|
||||
vector_store = get_vector_store()
|
||||
for (observer, observed), docs_with_embeddings in collection_embeddings.items():
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, observer, observed
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
observer,
|
||||
observed,
|
||||
)
|
||||
|
||||
# Build vector records with metadata for filtering
|
||||
vector_records: list[VectorRecord] = []
|
||||
doc_ids: list[str] = []
|
||||
for doc, embedding in docs_with_embeddings:
|
||||
doc_ids.append(doc.id)
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=doc.id,
|
||||
|
|
@ -567,21 +684,66 @@ async def create_observations(
|
|||
)
|
||||
)
|
||||
|
||||
# Retry vector upsert with exponential backoff
|
||||
# Retry vector upsert with exponential backoff (3 attempts)
|
||||
try:
|
||||
result = None
|
||||
async for attempt in AsyncRetrying(
|
||||
stop=stop_after_attempt(3),
|
||||
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
|
||||
retry=retry_if_exception_type(Exception)
|
||||
| retry_if_result(
|
||||
lambda res: res is not None and res.secondary_ok is False
|
||||
),
|
||||
reraise=True,
|
||||
):
|
||||
with attempt:
|
||||
await vector_store.upsert_many(namespace, vector_records)
|
||||
result = await vector_store.upsert_many(
|
||||
namespace, vector_records
|
||||
)
|
||||
|
||||
if result is not None and result.secondary_ok is False:
|
||||
# Partial success: primary has data but secondary doesn't
|
||||
# Keep as "pending" for reconciliation to sync secondary
|
||||
logger.warning(
|
||||
f"Partial sync for namespace {namespace}: {result.secondary_error}"
|
||||
)
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_attempts=models.Document.sync_attempts + 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
else:
|
||||
# Success: both primary and secondary stores have the data
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_state="synced",
|
||||
last_sync_at=func.now(),
|
||||
sync_attempts=0,
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
except Exception as e:
|
||||
# Final attempt failed - log but don't raise
|
||||
# Documents exist in DB, vectors can be added manually later
|
||||
# Total failure: primary write failed
|
||||
# Keep as "pending" for reconciliation to retry
|
||||
logger.error(
|
||||
f"Failed to upsert vectors for {namespace} after retries: {e}"
|
||||
f"Failed to upsert vectors for {namespace} after 3 retries: {e}"
|
||||
)
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_attempts=models.Document.sync_attempts + 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
except IntegrityError as e:
|
||||
await db.rollback()
|
||||
|
|
@ -652,8 +814,11 @@ async def is_rejected_duplicate(
|
|||
f"[DUPLICATE DETECTION] Deleting existing in favor of new. new='{doc.content}', existing='{existing_doc.content}'."
|
||||
)
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, observer, observed
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
observer,
|
||||
observed,
|
||||
)
|
||||
vector_deleted = False
|
||||
try:
|
||||
|
|
@ -685,9 +850,6 @@ async def cleanup_soft_deleted_documents(
|
|||
"""
|
||||
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)
|
||||
|
|
@ -729,8 +891,11 @@ async def cleanup_soft_deleted_documents(
|
|||
# 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
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
doc.workspace_name,
|
||||
doc.observer,
|
||||
doc.observed,
|
||||
)
|
||||
by_namespace.setdefault(namespace, []).append(doc.id)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,9 +2,15 @@ from logging import getLogger
|
|||
from typing import Any
|
||||
|
||||
from nanoid import generate as generate_nanoid
|
||||
from sqlalchemy import ColumnElement, Select, and_, func, select, text
|
||||
from sqlalchemy import ColumnElement, Select, and_, func, select, text, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from tenacity import AsyncRetrying, stop_after_attempt, wait_exponential
|
||||
from tenacity import (
|
||||
AsyncRetrying,
|
||||
retry_if_exception_type,
|
||||
retry_if_result,
|
||||
stop_after_attempt,
|
||||
wait_exponential,
|
||||
)
|
||||
|
||||
from src import models, schemas
|
||||
from src.config import settings
|
||||
|
|
@ -140,60 +146,155 @@ async def create_messages(
|
|||
|
||||
# Get vector store and namespace for this workspace's messages
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_message_namespace(workspace_name)
|
||||
namespace = vector_store.get_vector_namespace("message", workspace_name)
|
||||
|
||||
# Create MessageEmbedding entries and vector records
|
||||
# Create MessageEmbedding entries
|
||||
embedding_objects: list[models.MessageEmbedding] = []
|
||||
vector_records: list[VectorRecord] = []
|
||||
|
||||
# Check if pgvector is being used (primary or secondary)
|
||||
# If so, write embeddings to ORM since pgvector relies on postgres
|
||||
# Otherwise, store in memory for vector store upsert only
|
||||
pgvector_in_use = (
|
||||
settings.VECTOR_STORE.PRIMARY_TYPE == "pgvector"
|
||||
or settings.VECTOR_STORE.SECONDARY_TYPE == "pgvector"
|
||||
)
|
||||
|
||||
for message_obj in message_objects:
|
||||
embeddings = embedding_dict.get(message_obj.public_id, [])
|
||||
for chunk_index, embedding in enumerate(embeddings):
|
||||
# Create MessageEmbedding record for metadata tracking
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
message_id=message_obj.public_id,
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=message_obj.peer_name,
|
||||
chunk_index=chunk_index,
|
||||
)
|
||||
embedding_objects.append(embedding_obj)
|
||||
|
||||
# Create vector record for external vector store
|
||||
vector_id = f"{message_obj.public_id}_{chunk_index}"
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=vector_id,
|
||||
for embedding in embeddings:
|
||||
# Create MessageEmbedding record
|
||||
if pgvector_in_use:
|
||||
# pgvector in use: write embedding to ORM (postgres)
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
embedding=embedding,
|
||||
metadata={
|
||||
"message_id": message_obj.public_id,
|
||||
"session_name": session_name,
|
||||
"peer_name": message_obj.peer_name,
|
||||
"chunk_index": chunk_index,
|
||||
},
|
||||
message_id=message_obj.public_id,
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=message_obj.peer_name,
|
||||
)
|
||||
)
|
||||
else:
|
||||
# pgvector not in use: don't write embedding to postgres
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
message_id=message_obj.public_id,
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=message_obj.peer_name,
|
||||
)
|
||||
# Store embedding in memory for vector store upsert
|
||||
embedding_obj._pending_embedding = embedding
|
||||
embedding_obj.sync_state = "pending"
|
||||
embedding_objects.append(embedding_obj)
|
||||
|
||||
# Add all embedding metadata objects to the session
|
||||
if embedding_objects:
|
||||
db.add_all(embedding_objects)
|
||||
await db.flush()
|
||||
|
||||
# Track embedding IDs for sync state updates
|
||||
embedding_ids = [emb.id for emb in embedding_objects]
|
||||
|
||||
# Build vector records - source depends on whether pgvector is in use
|
||||
vector_records: list[VectorRecord] = []
|
||||
for emb in embedding_objects:
|
||||
if pgvector_in_use:
|
||||
# pgvector in use: embedding is on ORM object (numpy array)
|
||||
if emb.embedding is not None:
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=str(emb.id),
|
||||
embedding=[float(x) for x in emb.embedding],
|
||||
metadata={
|
||||
"message_id": emb.message_id,
|
||||
"session_name": emb.session_name,
|
||||
"peer_name": emb.peer_name,
|
||||
},
|
||||
)
|
||||
)
|
||||
else:
|
||||
# pgvector not in use: embedding is in _pending_embedding
|
||||
if (
|
||||
hasattr(emb, "_pending_embedding")
|
||||
and emb._pending_embedding is not None
|
||||
):
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=str(emb.id),
|
||||
embedding=list(emb._pending_embedding),
|
||||
metadata={
|
||||
"message_id": emb.message_id,
|
||||
"session_name": emb.session_name,
|
||||
"peer_name": emb.peer_name,
|
||||
},
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
# Upsert vectors to external vector store with retry
|
||||
if 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:
|
||||
# Final attempt failed - log but don't raise
|
||||
# MessageEmbedding records exist in DB, vectors can be added later
|
||||
logger.exception("Failed to upsert message vectors after retries")
|
||||
# Retry vector upsert with exponential backoff (3 attempts)
|
||||
if vector_records:
|
||||
try:
|
||||
result = None
|
||||
async for attempt in AsyncRetrying(
|
||||
stop=stop_after_attempt(3),
|
||||
wait=wait_exponential(multiplier=0.5, min=0.5, max=2.0),
|
||||
retry=retry_if_exception_type(Exception)
|
||||
| retry_if_result(
|
||||
lambda res: res is not None
|
||||
and res.secondary_ok is False
|
||||
),
|
||||
reraise=True,
|
||||
):
|
||||
with attempt:
|
||||
result = await vector_store.upsert_many(
|
||||
namespace, vector_records
|
||||
)
|
||||
|
||||
if result is not None and result.secondary_ok is False:
|
||||
# Partial success: primary has data but secondary doesn't
|
||||
# Keep as "pending" for reconciliation to sync secondary
|
||||
logger.warning(
|
||||
"Partial sync for message embeddings: %s",
|
||||
result.secondary_error,
|
||||
)
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id.in_(embedding_ids))
|
||||
.values(
|
||||
sync_attempts=models.MessageEmbedding.sync_attempts
|
||||
+ 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
else:
|
||||
# Success: both primary and secondary stores have the data
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id.in_(embedding_ids))
|
||||
.values(
|
||||
sync_state="synced",
|
||||
last_sync_at=func.now(),
|
||||
sync_attempts=0,
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
except Exception as e:
|
||||
# Total failure: primary write failed
|
||||
# Keep as "pending" for reconciliation to retry
|
||||
logger.error(
|
||||
f"Failed to upsert message vectors after 3 retries: {e}"
|
||||
)
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id.in_(embedding_ids))
|
||||
.values(
|
||||
sync_attempts=models.MessageEmbedding.sync_attempts + 1,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
except Exception:
|
||||
logger.exception(
|
||||
|
|
|
|||
|
|
@ -424,9 +424,7 @@ 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(
|
||||
select(models.MessageEmbedding.id).where(
|
||||
models.MessageEmbedding.session_name == session_name,
|
||||
models.MessageEmbedding.workspace_name == workspace_name,
|
||||
)
|
||||
|
|
@ -435,12 +433,11 @@ async def delete_session(
|
|||
vector_store = get_vector_store()
|
||||
|
||||
if embeddings:
|
||||
# Build vector IDs: {message_id}_{chunk_index}
|
||||
vector_ids = [f"{e.message_id}_{e.chunk_index}" for e in embeddings]
|
||||
vector_ids = [str(e.id) for e in embeddings]
|
||||
|
||||
# Try to delete from vector store (best effort)
|
||||
try:
|
||||
namespace = vector_store.get_message_namespace(workspace_name)
|
||||
namespace = vector_store.get_vector_namespace("message", workspace_name)
|
||||
await vector_store.delete_many(namespace, vector_ids)
|
||||
logger.debug(
|
||||
f"Deleted {len(vector_ids)} message vectors for session {session_name}"
|
||||
|
|
@ -480,8 +477,11 @@ async def delete_session(
|
|||
# Group document IDs by namespace (observer/observed)
|
||||
docs_by_namespace: dict[str, list[str]] = {}
|
||||
for doc in documents:
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, doc.observer, doc.observed
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
doc.observer,
|
||||
doc.observed,
|
||||
)
|
||||
docs_by_namespace.setdefault(namespace, []).append(doc.id)
|
||||
|
||||
|
|
|
|||
|
|
@ -326,7 +326,7 @@ async def delete_workspace(db: AsyncSession, workspace_name: str) -> schemas.Wor
|
|||
vector_store = get_vector_store()
|
||||
|
||||
# Delete message embeddings namespace for this workspace
|
||||
message_namespace = vector_store.get_message_namespace(workspace_name)
|
||||
message_namespace = vector_store.get_vector_namespace("message", workspace_name)
|
||||
try:
|
||||
await vector_store.delete_namespace(message_namespace)
|
||||
logger.debug(
|
||||
|
|
@ -343,8 +343,11 @@ async def delete_workspace(db: AsyncSession, workspace_name: str) -> schemas.Wor
|
|||
|
||||
# Delete document embeddings namespaces for each collection
|
||||
for collection in collections:
|
||||
doc_namespace = vector_store.get_document_namespace(
|
||||
workspace_name, collection.observer, collection.observed
|
||||
doc_namespace = vector_store.get_vector_namespace(
|
||||
"document",
|
||||
workspace_name,
|
||||
collection.observer,
|
||||
collection.observed,
|
||||
)
|
||||
try:
|
||||
await vector_store.delete_namespace(doc_namespace)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from fastapi import Depends
|
||||
from sqlalchemy import text
|
||||
|
|
@ -9,9 +8,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
|||
from src.config import settings
|
||||
from src.db import SessionLocal, request_context
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.vector_store import VectorStore
|
||||
|
||||
|
||||
async def get_db():
|
||||
"""FastAPI Dependency Generator for Database"""
|
||||
|
|
@ -65,17 +61,3 @@ async def tracked_db(operation_name: str | None = None):
|
|||
|
||||
|
||||
db: AsyncSession = Depends(get_db)
|
||||
|
||||
|
||||
def get_vector_store_dep() -> "VectorStore":
|
||||
"""FastAPI dependency for vector store.
|
||||
|
||||
This is a thin wrapper around get_vector_store() to allow for
|
||||
proper dependency injection in FastAPI routes.
|
||||
"""
|
||||
from src.vector_store import get_vector_store
|
||||
|
||||
return get_vector_store()
|
||||
|
||||
|
||||
vector_store: "VectorStore" = Depends(get_vector_store_dep)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from src.deriver.consumer import (
|
|||
process_item,
|
||||
process_representation_batch,
|
||||
)
|
||||
from src.deriver.vector_reconciliation import run_vector_reconciliation_cycle
|
||||
from src.dreamer.dream_scheduler import (
|
||||
DreamScheduler,
|
||||
get_dream_scheduler,
|
||||
|
|
@ -49,8 +50,8 @@ 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
|
||||
RECONCILIATION_INTERVAL_SECONDS = 300 # 5 minutes
|
||||
|
||||
|
||||
class QueueManager:
|
||||
|
|
@ -350,29 +351,22 @@ class QueueManager:
|
|||
"""
|
||||
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)
|
||||
- Reconciliation: every 5 minutes (sync vectors + clean up soft deletes)
|
||||
Only runs when pgvector is involved in a dual-store configuration
|
||||
"""
|
||||
# Track when each task should next run
|
||||
next_vector_cleanup = datetime.now(timezone.utc)
|
||||
next_queue_cleanup = datetime.now(timezone.utc)
|
||||
next_vector_reconciliation = (
|
||||
datetime.now(timezone.utc)
|
||||
if settings.VECTOR_STORE.should_run_reconciliation
|
||||
else None
|
||||
)
|
||||
|
||||
try:
|
||||
while not self.shutdown_event.is_set():
|
||||
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:
|
||||
|
|
@ -385,8 +379,28 @@ class QueueManager:
|
|||
seconds=QUEUE_CLEANUP_INTERVAL_SECONDS
|
||||
)
|
||||
|
||||
# Run vector store reconciliation if enabled and due
|
||||
if (
|
||||
next_vector_reconciliation is not None
|
||||
and now >= next_vector_reconciliation
|
||||
):
|
||||
try:
|
||||
logger.info("Running vector reconciliation cycle")
|
||||
await self._run_reconciliation()
|
||||
except Exception:
|
||||
logger.exception("Error during vector reconciliation")
|
||||
if settings.SENTRY.ENABLED:
|
||||
sentry_sdk.capture_exception()
|
||||
next_vector_reconciliation = now + timedelta(
|
||||
seconds=RECONCILIATION_INTERVAL_SECONDS
|
||||
)
|
||||
|
||||
# Sleep until next task is due or shutdown
|
||||
next_task_time = min(next_vector_cleanup, next_queue_cleanup)
|
||||
# Filter out None values when computing next task time
|
||||
task_times = [next_queue_cleanup]
|
||||
if next_vector_reconciliation is not None:
|
||||
task_times.append(next_vector_reconciliation)
|
||||
next_task_time = min(task_times)
|
||||
sleep_seconds = max(
|
||||
0, (next_task_time - datetime.now(timezone.utc)).total_seconds()
|
||||
)
|
||||
|
|
@ -405,26 +419,24 @@ class QueueManager:
|
|||
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 def _run_reconciliation(self) -> None:
|
||||
"""Run vector store reconciliation for sync + cleanup."""
|
||||
|
||||
async with tracked_db("vector_cleanup") as db:
|
||||
vector_store = get_vector_store()
|
||||
total_cleaned = 0
|
||||
metrics = await run_vector_reconciliation_cycle()
|
||||
|
||||
# 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"
|
||||
)
|
||||
if (
|
||||
metrics.total_synced > 0
|
||||
or metrics.total_failed > 0
|
||||
or metrics.total_cleaned > 0
|
||||
):
|
||||
logger.info(
|
||||
"Reconciliation: synced %s docs, %s message embeddings; failed %s docs, %s message embeddings; cleaned %s docs",
|
||||
metrics.documents_synced,
|
||||
metrics.message_embeddings_synced,
|
||||
metrics.documents_failed,
|
||||
metrics.message_embeddings_failed,
|
||||
metrics.documents_cleaned,
|
||||
)
|
||||
|
||||
async def _handle_processing_error(
|
||||
self,
|
||||
|
|
@ -854,6 +866,21 @@ class QueueManager:
|
|||
|
||||
async def main():
|
||||
logger.debug("Starting queue manager")
|
||||
|
||||
# Log reconciliation status
|
||||
if settings.VECTOR_STORE.should_run_reconciliation:
|
||||
logger.info(
|
||||
"Vector reconciliation: ENABLED (primary=%s, secondary=%s)",
|
||||
settings.VECTOR_STORE.PRIMARY_TYPE,
|
||||
settings.VECTOR_STORE.SECONDARY_TYPE,
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Vector reconciliation: DISABLED (primary=%s, secondary=%s)",
|
||||
settings.VECTOR_STORE.PRIMARY_TYPE,
|
||||
settings.VECTOR_STORE.SECONDARY_TYPE or "None",
|
||||
)
|
||||
|
||||
try:
|
||||
await init_cache()
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,388 @@
|
|||
"""
|
||||
Vector store reconciliation job.
|
||||
|
||||
This module provides a periodic reconciliation job that syncs documents and message
|
||||
embeddings to the vector store on a rolling basis, healing any missed writes.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import and_, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.sql.functions import func
|
||||
|
||||
from src import models
|
||||
from src.dependencies import tracked_db
|
||||
from src.vector_store import VectorRecord, VectorStore, get_vector_store
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Constants
|
||||
RECONCILIATION_BATCH_SIZE = 100
|
||||
RECONCILIATION_TIME_BUDGET_SECONDS = 240 # Leave headroom for other maintenance work
|
||||
MAX_SYNC_ATTEMPTS = 5 # After this many failures, mark as permanently_failed
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReconciliationMetrics:
|
||||
"""Metrics for a reconciliation cycle."""
|
||||
|
||||
documents_synced: int = 0
|
||||
documents_failed: int = 0
|
||||
documents_cleaned: int = 0
|
||||
message_embeddings_synced: int = 0
|
||||
message_embeddings_failed: int = 0
|
||||
|
||||
@property
|
||||
def total_synced(self) -> int:
|
||||
return self.documents_synced + self.message_embeddings_synced
|
||||
|
||||
@property
|
||||
def total_failed(self) -> int:
|
||||
return self.documents_failed + self.message_embeddings_failed
|
||||
|
||||
@property
|
||||
def total_cleaned(self) -> int:
|
||||
return self.documents_cleaned
|
||||
|
||||
|
||||
async def _get_documents_needing_sync(
|
||||
db: AsyncSession,
|
||||
batch_size: int = RECONCILIATION_BATCH_SIZE,
|
||||
) -> list[models.Document]:
|
||||
"""
|
||||
Get documents that need to be synced to the vector store.
|
||||
|
||||
Finds documents where:
|
||||
- not soft-deleted (deleted_at is NULL)
|
||||
- has an embedding stored in the database
|
||||
- sync_state is "pending" (never synced or retry needed)
|
||||
- Note: "synced" = done forever, "failed" = permanent failure (manual intervention)
|
||||
|
||||
Uses FOR UPDATE SKIP LOCKED to prevent concurrent processing.
|
||||
"""
|
||||
stmt = (
|
||||
select(models.Document)
|
||||
.where(
|
||||
and_(
|
||||
models.Document.deleted_at.is_(None),
|
||||
models.Document.embedding.isnot(None), # Must have embedding to sync
|
||||
models.Document.sync_state == "pending", # Only pending items
|
||||
)
|
||||
)
|
||||
.order_by(models.Document.last_sync_at.asc().nullsfirst())
|
||||
.limit(batch_size)
|
||||
.with_for_update(skip_locked=True)
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _get_message_embeddings_needing_sync(
|
||||
db: AsyncSession,
|
||||
batch_size: int = RECONCILIATION_BATCH_SIZE,
|
||||
) -> list[models.MessageEmbedding]:
|
||||
"""
|
||||
Get message embeddings that need to be synced to the vector store.
|
||||
|
||||
Finds embeddings where:
|
||||
- has an embedding stored in the database
|
||||
- sync_state is "pending" (never synced or retry needed)
|
||||
- Note: "synced" = done forever, "failed" = permanent failure (manual intervention)
|
||||
|
||||
Uses FOR UPDATE SKIP LOCKED to prevent concurrent processing.
|
||||
"""
|
||||
stmt = (
|
||||
select(models.MessageEmbedding)
|
||||
.where(
|
||||
and_(
|
||||
models.MessageEmbedding.embedding.isnot(
|
||||
None
|
||||
), # Must have embedding to sync
|
||||
models.MessageEmbedding.sync_state == "pending", # Only pending items
|
||||
)
|
||||
)
|
||||
.order_by(models.MessageEmbedding.last_sync_at.asc().nullsfirst())
|
||||
.limit(batch_size)
|
||||
.with_for_update(skip_locked=True)
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _sync_documents(
|
||||
db: AsyncSession,
|
||||
documents: list[models.Document],
|
||||
vector_store: VectorStore,
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
Sync a batch of documents to the vector store.
|
||||
|
||||
Returns (synced_count, failed_count).
|
||||
"""
|
||||
if not documents:
|
||||
return 0, 0
|
||||
|
||||
synced_count = 0
|
||||
failed_count = 0
|
||||
|
||||
# Group documents by namespace (workspace/observer/observed)
|
||||
by_namespace: dict[str, list[models.Document]] = {}
|
||||
for doc in documents:
|
||||
namespace = vector_store.get_vector_namespace(
|
||||
"document", doc.workspace_name, doc.observer, doc.observed
|
||||
)
|
||||
by_namespace.setdefault(namespace, []).append(doc)
|
||||
|
||||
# Sync each namespace batch
|
||||
for namespace, docs in by_namespace.items():
|
||||
doc_ids = [doc.id for doc in docs]
|
||||
|
||||
try:
|
||||
# Build vector records
|
||||
vector_records = [
|
||||
VectorRecord(
|
||||
id=doc.id,
|
||||
embedding=[float(x) for x in doc.embedding],
|
||||
metadata={
|
||||
"workspace_name": doc.workspace_name,
|
||||
"observer": doc.observer,
|
||||
"observed": doc.observed,
|
||||
"session_name": doc.session_name,
|
||||
"level": doc.level,
|
||||
},
|
||||
)
|
||||
for doc in docs
|
||||
if doc.embedding is not None
|
||||
]
|
||||
|
||||
result = None
|
||||
if vector_records:
|
||||
result = await vector_store.upsert_many(namespace, vector_records)
|
||||
|
||||
if result is not None and result.secondary_ok is False:
|
||||
logger.warning(
|
||||
"Partial sync for namespace %s: %s",
|
||||
namespace,
|
||||
result.secondary_error,
|
||||
)
|
||||
# Increment attempts and mark as failed if we've hit max attempts
|
||||
for doc in docs:
|
||||
new_attempts = doc.sync_attempts + 1
|
||||
new_state = (
|
||||
"failed" if new_attempts >= MAX_SYNC_ATTEMPTS else "pending"
|
||||
)
|
||||
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id == doc.id)
|
||||
.values(
|
||||
sync_state=new_state,
|
||||
sync_attempts=new_attempts,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
failed_count += len(docs)
|
||||
continue
|
||||
|
||||
# Mark as synced
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(doc_ids))
|
||||
.values(
|
||||
sync_state="synced",
|
||||
last_sync_at=func.now(),
|
||||
sync_attempts=0,
|
||||
)
|
||||
)
|
||||
synced_count += len(docs)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to sync documents to {namespace}: {e}")
|
||||
# Increment attempts and mark as failed if we've hit max attempts
|
||||
for doc in docs:
|
||||
new_attempts = doc.sync_attempts + 1
|
||||
new_state = "failed" if new_attempts >= MAX_SYNC_ATTEMPTS else "pending"
|
||||
|
||||
await db.execute(
|
||||
update(models.Document)
|
||||
.where(models.Document.id == doc.id)
|
||||
.values(
|
||||
sync_state=new_state,
|
||||
sync_attempts=new_attempts,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
failed_count += len(docs)
|
||||
|
||||
return synced_count, failed_count
|
||||
|
||||
|
||||
async def _sync_message_embeddings(
|
||||
db: AsyncSession,
|
||||
embeddings: list[models.MessageEmbedding],
|
||||
vector_store: VectorStore,
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
Sync a batch of message embeddings to the vector store.
|
||||
|
||||
Returns (synced_count, failed_count).
|
||||
"""
|
||||
if not embeddings:
|
||||
return 0, 0
|
||||
|
||||
synced_count = 0
|
||||
failed_count = 0
|
||||
|
||||
# Group by namespace (workspace)
|
||||
by_namespace: dict[str, list[models.MessageEmbedding]] = {}
|
||||
for emb in embeddings:
|
||||
namespace = vector_store.get_vector_namespace("message", emb.workspace_name)
|
||||
by_namespace.setdefault(namespace, []).append(emb)
|
||||
|
||||
# Sync each namespace batch
|
||||
for namespace, embs in by_namespace.items():
|
||||
emb_ids = [emb.id for emb in embs]
|
||||
|
||||
try:
|
||||
# Build vector records
|
||||
vector_records = [
|
||||
VectorRecord(
|
||||
id=str(emb.id),
|
||||
embedding=[float(x) for x in emb.embedding],
|
||||
metadata={
|
||||
"message_id": emb.message_id,
|
||||
"session_name": emb.session_name,
|
||||
"peer_name": emb.peer_name,
|
||||
},
|
||||
)
|
||||
for emb in embs
|
||||
if emb.embedding is not None
|
||||
]
|
||||
|
||||
result = None
|
||||
if vector_records:
|
||||
result = await vector_store.upsert_many(namespace, vector_records)
|
||||
|
||||
if result is not None and result.secondary_ok is False:
|
||||
logger.warning(
|
||||
"Partial sync for namespace %s: %s",
|
||||
namespace,
|
||||
result.secondary_error,
|
||||
)
|
||||
# Increment attempts and mark as failed if we've hit max attempts
|
||||
for emb in embs:
|
||||
new_attempts = emb.sync_attempts + 1
|
||||
new_state = (
|
||||
"failed" if new_attempts >= MAX_SYNC_ATTEMPTS else "pending"
|
||||
)
|
||||
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id == emb.id)
|
||||
.values(
|
||||
sync_state=new_state,
|
||||
sync_attempts=new_attempts,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
failed_count += len(embs)
|
||||
continue
|
||||
|
||||
# Mark as synced
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id.in_(emb_ids))
|
||||
.values(
|
||||
sync_state="synced",
|
||||
last_sync_at=func.now(),
|
||||
sync_attempts=0,
|
||||
)
|
||||
)
|
||||
synced_count += len(embs)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to sync message embeddings to {namespace}: {e}")
|
||||
# Increment attempts and mark as failed if we've hit max attempts
|
||||
for emb in embs:
|
||||
new_attempts = emb.sync_attempts + 1
|
||||
new_state = "failed" if new_attempts >= MAX_SYNC_ATTEMPTS else "pending"
|
||||
|
||||
await db.execute(
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id == emb.id)
|
||||
.values(
|
||||
sync_state=new_state,
|
||||
sync_attempts=new_attempts,
|
||||
last_sync_at=func.now(),
|
||||
)
|
||||
)
|
||||
failed_count += len(embs)
|
||||
|
||||
return synced_count, failed_count
|
||||
|
||||
|
||||
async def run_vector_reconciliation_cycle() -> ReconciliationMetrics:
|
||||
"""
|
||||
Run a complete reconciliation cycle.
|
||||
|
||||
Runs a rolling sweep to reconcile missing vectors and clean up soft deletes.
|
||||
Uses batching and FOR UPDATE SKIP LOCKED for safe concurrent operation.
|
||||
|
||||
Returns metrics about what was synced.
|
||||
"""
|
||||
metrics = ReconciliationMetrics()
|
||||
vector_store = get_vector_store()
|
||||
deadline = time.monotonic() + RECONCILIATION_TIME_BUDGET_SECONDS
|
||||
|
||||
from src.crud.document import cleanup_soft_deleted_documents
|
||||
|
||||
print("Running vector reconciliation cycle")
|
||||
async with tracked_db("reconciliation") as db:
|
||||
while time.monotonic() < deadline:
|
||||
did_work = False
|
||||
|
||||
# Reconcile documents
|
||||
docs = await _get_documents_needing_sync(db)
|
||||
if docs:
|
||||
synced, failed = await _sync_documents(db, docs, vector_store)
|
||||
metrics.documents_synced += synced
|
||||
metrics.documents_failed += failed
|
||||
await db.commit()
|
||||
did_work = True
|
||||
|
||||
if time.monotonic() >= deadline:
|
||||
break
|
||||
|
||||
# Reconcile message embeddings
|
||||
embs = await _get_message_embeddings_needing_sync(db)
|
||||
if embs:
|
||||
synced, failed = await _sync_message_embeddings(db, embs, vector_store)
|
||||
metrics.message_embeddings_synced += synced
|
||||
metrics.message_embeddings_failed += failed
|
||||
await db.commit()
|
||||
did_work = True
|
||||
|
||||
if time.monotonic() >= deadline:
|
||||
break
|
||||
|
||||
# Clean up soft-deleted documents
|
||||
cleaned = await cleanup_soft_deleted_documents(
|
||||
db,
|
||||
vector_store,
|
||||
batch_size=RECONCILIATION_BATCH_SIZE,
|
||||
)
|
||||
if cleaned:
|
||||
metrics.documents_cleaned += cleaned
|
||||
did_work = True
|
||||
|
||||
if not did_work:
|
||||
print("No work done, breaking")
|
||||
break
|
||||
print("Vector reconciliation cycle completed")
|
||||
|
||||
return metrics
|
||||
|
|
@ -109,6 +109,26 @@ class FileProcessingError(HonchoException):
|
|||
detail = "File processing error"
|
||||
|
||||
|
||||
@final
|
||||
class PartialVectorSyncException(HonchoException):
|
||||
"""
|
||||
Exception raised when vector upsert partially succeeds.
|
||||
|
||||
This indicates the primary store succeeded but secondary store failed.
|
||||
The data is queryable but not fully replicated.
|
||||
"""
|
||||
|
||||
status_code = 500
|
||||
detail = "Vector partially synced to primary store only"
|
||||
|
||||
def __init__(self, primary_success: bool, secondary_error: Exception):
|
||||
self.primary_success = primary_success
|
||||
self.secondary_error = secondary_error
|
||||
super().__init__(
|
||||
f"Partial sync: primary={'succeeded' if primary_success else 'failed'}, secondary failed with: {secondary_error}"
|
||||
)
|
||||
|
||||
|
||||
class LLMError(Exception):
|
||||
"""Exception raised when an LLM call fails.
|
||||
|
||||
|
|
|
|||
|
|
@ -125,6 +125,10 @@ async def lifespan(_: FastAPI):
|
|||
try:
|
||||
yield
|
||||
finally:
|
||||
# Import here to avoid circular import at module load time
|
||||
from src.vector_store import close_vector_store
|
||||
|
||||
await close_vector_store()
|
||||
await close_cache()
|
||||
await engine.dispose()
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Any, final
|
|||
|
||||
from dotenv import load_dotenv
|
||||
from nanoid import generate as generate_nanoid
|
||||
from pgvector.sqlalchemy import Vector
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
|
|
@ -20,11 +21,11 @@ from sqlalchemy import (
|
|||
text,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB, TEXT
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.orm import Mapped, MappedColumn, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
from typing_extensions import override
|
||||
|
||||
from src.utils.types import DocumentLevel, TaskType
|
||||
from src.utils.types import DocumentLevel, TaskType, VectorSyncState
|
||||
|
||||
from .db import Base
|
||||
|
||||
|
|
@ -271,21 +272,13 @@ class Message(Base):
|
|||
|
||||
@final
|
||||
class MessageEmbedding(Base):
|
||||
"""
|
||||
Stores metadata for message embeddings.
|
||||
|
||||
Note: The actual embedding vectors are stored in the external vector store
|
||||
(Turbopuffer or LanceDB), not in PostgreSQL. This table maintains the
|
||||
relationship between messages and their embeddings, along with metadata
|
||||
needed for filtering and lookups.
|
||||
"""
|
||||
|
||||
__tablename__: str = "message_embeddings"
|
||||
|
||||
id: Mapped[int] = mapped_column(
|
||||
BigInteger, Identity(), primary_key=True, autoincrement=True
|
||||
)
|
||||
content: Mapped[str] = mapped_column(TEXT)
|
||||
embedding: MappedColumn[Any] = mapped_column(Vector(1536), nullable=True)
|
||||
message_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("messages.public_id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
|
|
@ -297,8 +290,16 @@ class MessageEmbedding(Base):
|
|||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), index=True
|
||||
)
|
||||
# Chunk index for messages that are split into multiple embeddings
|
||||
chunk_index: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
# Vector sync state tracking
|
||||
sync_state: Mapped[VectorSyncState] = mapped_column(
|
||||
TEXT, nullable=False, server_default="pending", index=True
|
||||
)
|
||||
last_sync_at: Mapped[datetime.datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
sync_attempts: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, default=0, server_default=text("0")
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
# Compound foreign key constraints
|
||||
|
|
@ -310,6 +311,14 @@ class MessageEmbedding(Base):
|
|||
["peer_name", "workspace_name"],
|
||||
["peers.name", "peers.workspace_name"],
|
||||
),
|
||||
# HNSW index on embedding column for efficient similarity search
|
||||
Index(
|
||||
"ix_message_embeddings_embedding_hnsw",
|
||||
"embedding",
|
||||
postgresql_using="hnsw",
|
||||
postgresql_with={"m": 16, "ef_construction": 64},
|
||||
postgresql_ops={"embedding": "vector_cosine_ops"},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -359,14 +368,6 @@ class Collection(Base):
|
|||
|
||||
@final
|
||||
class Document(Base):
|
||||
"""
|
||||
Stores document metadata and content.
|
||||
|
||||
Note: The actual embedding vectors are stored in the external vector store
|
||||
(Turbopuffer or LanceDB), not in PostgreSQL. The vector ID is the document's
|
||||
primary key (id field).
|
||||
"""
|
||||
|
||||
__tablename__: str = "documents"
|
||||
id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True)
|
||||
internal_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||
|
|
@ -379,6 +380,7 @@ class Document(Base):
|
|||
times_derived: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, server_default=text("1")
|
||||
)
|
||||
embedding: MappedColumn[Any] = mapped_column(Vector(1536), nullable=True)
|
||||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), index=True
|
||||
)
|
||||
|
|
@ -392,6 +394,18 @@ class Document(Base):
|
|||
deleted_at: Mapped[datetime.datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True, index=True, default=None
|
||||
)
|
||||
|
||||
# Vector sync state tracking
|
||||
sync_state: Mapped[VectorSyncState] = mapped_column(
|
||||
TEXT, nullable=False, server_default="pending", index=True
|
||||
)
|
||||
last_sync_at: Mapped[datetime.datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
sync_attempts: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, default=0, server_default=text("0")
|
||||
)
|
||||
|
||||
collection = relationship("Collection", back_populates="documents")
|
||||
|
||||
__table_args__ = (
|
||||
|
|
@ -422,6 +436,16 @@ class Document(Base):
|
|||
["session_name", "workspace_name"],
|
||||
["sessions.name", "sessions.workspace_name"],
|
||||
),
|
||||
# HNSW index on embedding column
|
||||
Index(
|
||||
"ix_documents_embedding_hnsw",
|
||||
"embedding",
|
||||
postgresql_using="hnsw", # HNSW index type
|
||||
postgresql_with={"m": 16, "ef_construction": 64}, # HNSW parameters
|
||||
postgresql_ops={
|
||||
"embedding": "vector_cosine_ops"
|
||||
}, # Cosine distance operator
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ async def _semantic_search(
|
|||
|
||||
# Get vector store and namespace for this workspace's messages
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_message_namespace(workspace_name)
|
||||
namespace = vector_store.get_vector_namespace("message", workspace_name)
|
||||
|
||||
# Build vector store filters from the provided filters
|
||||
vector_filters: dict[str, Any] = {}
|
||||
|
|
@ -116,17 +116,14 @@ async def _semantic_search(
|
|||
if not vector_results:
|
||||
return []
|
||||
|
||||
# Extract message IDs from vector results (vector ID format: {message_public_id}_{chunk_index})
|
||||
# Extract message IDs from vector metadata
|
||||
# Use dict to deduplicate while preserving order (dict keys maintain insertion order in Python 3.7+)
|
||||
seen_message_ids: dict[str, None] = {}
|
||||
|
||||
for result in vector_results:
|
||||
# Vector ID format: {message_public_id}_{chunk_index}
|
||||
parts = result.id.rsplit("_", 1)
|
||||
if len(parts) >= 1:
|
||||
message_id = parts[0]
|
||||
if message_id not in seen_message_ids:
|
||||
seen_message_ids[message_id] = None
|
||||
message_id = result.metadata.get("message_id")
|
||||
if message_id and message_id not in seen_message_ids:
|
||||
seen_message_ids[message_id] = None
|
||||
|
||||
message_ids = list(seen_message_ids.keys())
|
||||
|
||||
|
|
|
|||
|
|
@ -3,3 +3,4 @@ from typing import Literal
|
|||
SupportedProviders = Literal["anthropic", "openai", "google", "groq", "custom", "vllm"]
|
||||
TaskType = Literal["webhook", "summary", "representation", "dream", "deletion"]
|
||||
DocumentLevel = Literal["explicit", "deductive"]
|
||||
VectorSyncState = Literal["synced", "pending", "failed"]
|
||||
|
|
|
|||
|
|
@ -3,28 +3,46 @@ Vector store abstraction layer for Honcho.
|
|||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from functools import cache
|
||||
from typing import Any, ClassVar, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from src.config import settings
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorRecord:
|
||||
class VectorRecord(BaseModel):
|
||||
"""A single vector record to be stored in the vector store."""
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
id: str
|
||||
embedding: list[float]
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QueryResult:
|
||||
class VectorQueryResult(BaseModel):
|
||||
"""A single result from a vector query."""
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
id: str
|
||||
score: float # Distance/similarity score (lower = more similar for cosine distance)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class VectorUpsertResult(BaseModel):
|
||||
"""Result for a vector upsert operation."""
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(
|
||||
extra="forbid",
|
||||
frozen=True,
|
||||
arbitrary_types_allowed=True,
|
||||
)
|
||||
|
||||
primary_ok: bool
|
||||
secondary_ok: bool | None = None
|
||||
secondary_error: Exception | None = None
|
||||
|
||||
|
||||
class VectorStore(ABC):
|
||||
|
|
@ -49,62 +67,52 @@ class VectorStore(ABC):
|
|||
self.namespace_prefix = settings.VECTOR_STORE.NAMESPACE
|
||||
|
||||
# === Namespace helpers ===
|
||||
def get_document_namespace(
|
||||
self, workspace_name: str, observer: str, observed: str
|
||||
def get_vector_namespace(
|
||||
self,
|
||||
namespace_type: Literal["document", "message"],
|
||||
workspace_name: str,
|
||||
observer: str | None = None,
|
||||
observed: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the namespace for document embeddings (per collection).
|
||||
Get the namespace for document or message embeddings.
|
||||
|
||||
Args:
|
||||
namespace_type: "document" or "message"
|
||||
workspace_name: Name of the workspace
|
||||
observer: Name of the observing peer
|
||||
observed: Name of the observed peer
|
||||
observer: Name of the observing peer (document only)
|
||||
observed: Name of the observed peer (document only)
|
||||
|
||||
Returns:
|
||||
Namespace string in format: {prefix}.{workspace}.{observer}.{observed}
|
||||
Namespace string in format:
|
||||
- document: {prefix}.{workspace}.{observer}.{observed}
|
||||
- message: {prefix}.{workspace}.messages
|
||||
"""
|
||||
return f"{self.namespace_prefix}.{workspace_name}.{observer}.{observed}"
|
||||
|
||||
def get_message_namespace(self, workspace_name: str) -> str:
|
||||
"""
|
||||
Get the namespace for message embeddings (per workspace).
|
||||
|
||||
Args:
|
||||
workspace_name: Name of the workspace
|
||||
|
||||
Returns:
|
||||
Namespace string in format: {prefix}.{workspace}.messages
|
||||
"""
|
||||
return f"{self.namespace_prefix}.{workspace_name}.messages"
|
||||
if namespace_type == "document":
|
||||
if observer is None or observed is None:
|
||||
raise ValueError(
|
||||
"observer and observed are required for document namespaces"
|
||||
)
|
||||
return f"{self.namespace_prefix}.{workspace_name}.{observer}.{observed}"
|
||||
if namespace_type == "message":
|
||||
return f"{self.namespace_prefix}.{workspace_name}.messages"
|
||||
|
||||
# === Core operations ===
|
||||
@abstractmethod
|
||||
async def upsert(
|
||||
self,
|
||||
namespace: str,
|
||||
vector: VectorRecord,
|
||||
) -> None:
|
||||
"""
|
||||
Upsert a single vector into the store.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vector in
|
||||
vector: VectorRecord containing id, embedding, and optional metadata
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> None:
|
||||
) -> VectorUpsertResult:
|
||||
"""
|
||||
Upsert multiple vectors into the store.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vectors in
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
|
||||
Returns:
|
||||
Result describing primary/secondary outcomes.
|
||||
"""
|
||||
...
|
||||
|
||||
|
|
@ -117,7 +125,7 @@ class VectorStore(ABC):
|
|||
top_k: int = 10,
|
||||
filters: dict[str, Any] | None = None,
|
||||
max_distance: float | None = None,
|
||||
) -> list[QueryResult]:
|
||||
) -> list[VectorQueryResult]:
|
||||
"""
|
||||
Query for similar vectors.
|
||||
|
||||
|
|
@ -129,7 +137,7 @@ class VectorStore(ABC):
|
|||
max_distance: Optional maximum distance threshold (cosine distance)
|
||||
|
||||
Returns:
|
||||
List of QueryResult objects, ordered by similarity (most similar first)
|
||||
List of VectorQueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
...
|
||||
|
||||
|
|
@ -154,43 +162,65 @@ class VectorStore(ABC):
|
|||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def close(self) -> None:
|
||||
"""
|
||||
Close any open connections and release resources.
|
||||
|
||||
# Singleton instance
|
||||
_vector_store_instance: VectorStore | None = None
|
||||
Subclasses should override this if they maintain persistent connections.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
# Import implementations after base classes are defined to avoid circular imports
|
||||
from src.vector_store.composite import CompositeVectorStore # noqa: E402
|
||||
from src.vector_store.lancedb import LanceDBVectorStore # noqa: E402
|
||||
from src.vector_store.pgvector import PgVectorStore # noqa: E402
|
||||
from src.vector_store.turbopuffer import TurbopufferVectorStore # noqa: E402
|
||||
|
||||
|
||||
def _create_store_by_type(store_type: str) -> VectorStore:
|
||||
"""Create a vector store instance by type name."""
|
||||
if store_type == "turbopuffer":
|
||||
return TurbopufferVectorStore()
|
||||
elif store_type == "lancedb":
|
||||
return LanceDBVectorStore()
|
||||
elif store_type == "pgvector":
|
||||
return PgVectorStore()
|
||||
else:
|
||||
raise ValueError(f"Unknown vector store type: {store_type}")
|
||||
|
||||
|
||||
def _create_vector_store() -> VectorStore:
|
||||
"""
|
||||
Create a new vector store instance based on configuration.
|
||||
|
||||
If SECONDARY_TYPE is set, returns a CompositeVectorStore that:
|
||||
- Writes to both primary and secondary stores
|
||||
- Reads from primary only, falls back to secondary on failure
|
||||
|
||||
Returns:
|
||||
The vector store instance based on configuration.
|
||||
|
||||
Raises:
|
||||
ValueError: If the configured vector store type is invalid.
|
||||
"""
|
||||
store_type = settings.VECTOR_STORE.TYPE
|
||||
primary = _create_store_by_type(settings.VECTOR_STORE.PRIMARY_TYPE)
|
||||
|
||||
if store_type == "turbopuffer":
|
||||
from src.vector_store.turbopuffer import TurbopufferVectorStore
|
||||
if settings.VECTOR_STORE.SECONDARY_TYPE:
|
||||
secondary = _create_store_by_type(settings.VECTOR_STORE.SECONDARY_TYPE)
|
||||
return CompositeVectorStore(primary=primary, secondary=secondary)
|
||||
|
||||
return TurbopufferVectorStore()
|
||||
elif store_type == "lancedb":
|
||||
from src.vector_store.lancedb import LanceDBVectorStore
|
||||
|
||||
return LanceDBVectorStore()
|
||||
else:
|
||||
raise ValueError(f"Unknown vector store type: {store_type}")
|
||||
return primary
|
||||
|
||||
|
||||
@cache
|
||||
def get_vector_store() -> VectorStore:
|
||||
"""
|
||||
FastAPI dependency that provides the configured vector store instance (singleton).
|
||||
Get the configured vector store instance (singleton).
|
||||
|
||||
This function is designed to be used as a FastAPI dependency:
|
||||
vector_store: VectorStore = Depends(get_vector_store)
|
||||
|
||||
It can also be called directly for non-request contexts (e.g., background tasks).
|
||||
Uses functools.cache to ensure only one instance is created per process.
|
||||
This is asyncio-safe since there are no await points in the creation path.
|
||||
|
||||
Returns:
|
||||
The vector store instance based on configuration.
|
||||
|
|
@ -198,28 +228,32 @@ def get_vector_store() -> VectorStore:
|
|||
Raises:
|
||||
ValueError: If the configured vector store type is invalid.
|
||||
"""
|
||||
global _vector_store_instance
|
||||
|
||||
if _vector_store_instance is None:
|
||||
_vector_store_instance = _create_vector_store()
|
||||
|
||||
return _vector_store_instance
|
||||
return _create_vector_store()
|
||||
|
||||
|
||||
def reset_vector_store() -> None:
|
||||
async def close_vector_store() -> None:
|
||||
"""
|
||||
Reset the vector store singleton instance.
|
||||
Close the vector store and release resources.
|
||||
|
||||
This is primarily useful for testing to ensure a fresh instance is created.
|
||||
Call this during application shutdown to cleanly close connections.
|
||||
After calling this, you must call get_vector_store.cache_clear() if you
|
||||
want to create a new instance.
|
||||
"""
|
||||
global _vector_store_instance
|
||||
_vector_store_instance = None
|
||||
# Check if an instance was ever created
|
||||
if (
|
||||
get_vector_store.cache_info().hits > 0
|
||||
or get_vector_store.cache_info().misses > 0
|
||||
):
|
||||
store = get_vector_store()
|
||||
await store.close()
|
||||
get_vector_store.cache_clear()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"VectorStore",
|
||||
"VectorRecord",
|
||||
"QueryResult",
|
||||
"VectorQueryResult",
|
||||
"VectorUpsertResult",
|
||||
"get_vector_store",
|
||||
"reset_vector_store",
|
||||
"close_vector_store",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,261 @@
|
|||
"""
|
||||
Composite vector store implementation.
|
||||
|
||||
This module provides a composite VectorStore that writes to two stores
|
||||
and reads from primary with fallback to secondary.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CompositeVectorStore(VectorStore):
|
||||
"""
|
||||
Composite vector store with dual-write and fallback-read.
|
||||
|
||||
Behavior:
|
||||
- Writes go to BOTH primary and secondary stores
|
||||
- Reads try primary only, fall back to secondary on failure
|
||||
|
||||
Migration Strategy:
|
||||
- Primary = source of truth (pgvector)
|
||||
- Secondary = target being populated (turbopuffer)
|
||||
- Reconciliation job syncs data from primary to secondary via dual-writes
|
||||
- Reads use primary, fallback to secondary only on exception (network errors, etc.)
|
||||
- After migration completes, remove secondary from config
|
||||
"""
|
||||
|
||||
primary: VectorStore
|
||||
secondary: VectorStore
|
||||
|
||||
def __init__(self, primary: VectorStore, secondary: VectorStore):
|
||||
"""
|
||||
Initialize the composite vector store.
|
||||
|
||||
Args:
|
||||
primary: The primary vector store (reads prefer this)
|
||||
secondary: The secondary vector store (fallback for reads)
|
||||
"""
|
||||
super().__init__()
|
||||
self.primary = primary
|
||||
self.secondary = secondary
|
||||
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> VectorUpsertResult:
|
||||
"""
|
||||
Upsert multiple vectors to both stores.
|
||||
|
||||
Success cases:
|
||||
- Primary ✓, Secondary ✓ → Success (fully synced)
|
||||
- Primary ✓, Secondary ✗ → Returns partial result (not fully synced)
|
||||
|
||||
Failure cases:
|
||||
- Primary ✗, Secondary ✓ → Raises primary exception (weird, shouldn't happen)
|
||||
- Primary ✗, Secondary ✗ → Raises primary exception (total failure)
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vectors in
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
|
||||
Returns:
|
||||
Result describing primary/secondary outcomes.
|
||||
|
||||
Raises:
|
||||
Exception: Primary failed (secondary state doesn't matter)
|
||||
"""
|
||||
if not vectors:
|
||||
return VectorUpsertResult(primary_ok=True, secondary_ok=True)
|
||||
|
||||
# Write to both stores concurrently
|
||||
primary_task = asyncio.create_task(self.primary.upsert_many(namespace, vectors))
|
||||
secondary_task = asyncio.create_task(
|
||||
self.secondary.upsert_many(namespace, vectors)
|
||||
)
|
||||
|
||||
# Wait for both, gathering exceptions
|
||||
results = await asyncio.gather(
|
||||
primary_task, secondary_task, return_exceptions=True
|
||||
)
|
||||
|
||||
primary_result, secondary_result = results
|
||||
|
||||
# Case 1: Both failed → raise primary exception
|
||||
if isinstance(primary_result, Exception) and isinstance(
|
||||
secondary_result, Exception
|
||||
):
|
||||
logger.error(
|
||||
f"Both primary and secondary upsert failed for namespace {namespace}. Primary: {primary_result}, Secondary: {secondary_result}"
|
||||
)
|
||||
raise primary_result
|
||||
|
||||
# Case 2: Primary failed, secondary succeeded → raise primary exception (weird case)
|
||||
if isinstance(primary_result, Exception):
|
||||
logger.error(
|
||||
f"Primary upsert failed but secondary succeeded for namespace {namespace}: {primary_result}"
|
||||
)
|
||||
raise primary_result
|
||||
|
||||
# Case 3: Primary succeeded, secondary failed → return partial result
|
||||
if isinstance(secondary_result, Exception):
|
||||
logger.warning(
|
||||
f"Primary upsert succeeded but secondary failed for namespace {namespace}: {secondary_result}"
|
||||
)
|
||||
return VectorUpsertResult(
|
||||
primary_ok=True,
|
||||
secondary_ok=False,
|
||||
secondary_error=secondary_result,
|
||||
)
|
||||
|
||||
# Case 4: Both succeeded → log success
|
||||
logger.debug(
|
||||
f"Dual-write upserted {len(vectors)} vectors to namespace {namespace}"
|
||||
)
|
||||
return VectorUpsertResult(primary_ok=True, secondary_ok=True)
|
||||
|
||||
async def query(
|
||||
self,
|
||||
namespace: str,
|
||||
embedding: list[float],
|
||||
*,
|
||||
top_k: int = 10,
|
||||
filters: dict[str, Any] | None = None,
|
||||
max_distance: float | None = None,
|
||||
) -> list[VectorQueryResult]:
|
||||
"""
|
||||
Query for similar vectors, trying primary first then falling back to secondary on failure.
|
||||
|
||||
Primary is the source of truth. Secondary is only used if primary query raises an
|
||||
exception (network errors, timeouts, etc.). If primary returns empty results ([]),
|
||||
that's considered success and we don't query secondary.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to query
|
||||
embedding: The query embedding vector
|
||||
top_k: Maximum number of results to return
|
||||
filters: Optional metadata filters
|
||||
max_distance: Optional maximum distance threshold (cosine distance)
|
||||
|
||||
Returns:
|
||||
List of VectorQueryResult objects, ordered by similarity (most similar first)
|
||||
|
||||
Raises:
|
||||
Exception: If both primary and secondary queries fail
|
||||
"""
|
||||
try:
|
||||
results = await self.primary.query(
|
||||
namespace,
|
||||
embedding,
|
||||
top_k=top_k,
|
||||
filters=filters,
|
||||
max_distance=max_distance,
|
||||
)
|
||||
logger.debug(
|
||||
f"Primary query returned {len(results)} results from namespace {namespace}"
|
||||
)
|
||||
return results
|
||||
except Exception as primary_error:
|
||||
logger.warning(
|
||||
f"Primary query failed for namespace {namespace}: {primary_error}, attempting fallback to secondary"
|
||||
)
|
||||
try:
|
||||
results = await self.secondary.query(
|
||||
namespace,
|
||||
embedding,
|
||||
top_k=top_k,
|
||||
filters=filters,
|
||||
max_distance=max_distance,
|
||||
)
|
||||
logger.warning(
|
||||
f"Secondary query returned {len(results)} results from namespace {namespace}"
|
||||
)
|
||||
return results
|
||||
except Exception as secondary_error:
|
||||
logger.error(
|
||||
f"Both primary and secondary queries failed for namespace {namespace}. Primary: {primary_error}, Secondary: {secondary_error}"
|
||||
)
|
||||
raise primary_error from secondary_error
|
||||
|
||||
async def delete_many(self, namespace: str, ids: list[str]) -> None:
|
||||
"""
|
||||
Delete vectors from both stores.
|
||||
|
||||
Args:
|
||||
namespace: The namespace containing the vectors
|
||||
ids: List of vector identifiers to delete
|
||||
"""
|
||||
if not ids:
|
||||
return
|
||||
|
||||
# Delete from both stores concurrently
|
||||
primary_task = asyncio.create_task(self.primary.delete_many(namespace, ids))
|
||||
secondary_task = asyncio.create_task(self.secondary.delete_many(namespace, ids))
|
||||
|
||||
results = await asyncio.gather(
|
||||
primary_task, secondary_task, return_exceptions=True
|
||||
)
|
||||
|
||||
primary_result, secondary_result = results
|
||||
|
||||
# Primary failure is critical - raise it
|
||||
if isinstance(primary_result, Exception):
|
||||
logger.error(
|
||||
f"Primary vector store delete failed for namespace {namespace}: {primary_result}"
|
||||
)
|
||||
raise primary_result
|
||||
|
||||
if isinstance(secondary_result, Exception):
|
||||
logger.warning(
|
||||
f"Secondary vector store delete failed for namespace {namespace}: {secondary_result}"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"Dual-delete removed {len(ids)} vectors from namespace {namespace}"
|
||||
)
|
||||
|
||||
async def delete_namespace(self, namespace: str) -> None:
|
||||
"""
|
||||
Delete an entire namespace from both stores.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to delete
|
||||
"""
|
||||
# Delete from both stores concurrently
|
||||
primary_task = asyncio.create_task(self.primary.delete_namespace(namespace))
|
||||
secondary_task = asyncio.create_task(self.secondary.delete_namespace(namespace))
|
||||
|
||||
results = await asyncio.gather(
|
||||
primary_task, secondary_task, return_exceptions=True
|
||||
)
|
||||
|
||||
primary_result, secondary_result = results
|
||||
|
||||
# Primary failure is critical - raise it
|
||||
if isinstance(primary_result, Exception):
|
||||
logger.error(
|
||||
f"Primary vector store namespace delete failed for {namespace}: {primary_result}"
|
||||
)
|
||||
raise primary_result
|
||||
|
||||
if isinstance(secondary_result, Exception):
|
||||
logger.warning(
|
||||
f"Secondary vector store namespace delete failed for {namespace}: {secondary_result}"
|
||||
)
|
||||
|
||||
logger.debug(f"Dual-delete removed namespace {namespace}")
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close both primary and secondary vector stores."""
|
||||
await asyncio.gather(
|
||||
self.primary.close(),
|
||||
self.secondary.close(),
|
||||
return_exceptions=True,
|
||||
)
|
||||
logger.debug("Composite vector store closed")
|
||||
|
|
@ -5,6 +5,7 @@ This module provides a LanceDB-based implementation of the VectorStore interface
|
|||
for use in self-hosted deployments of Honcho.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any, cast
|
||||
|
||||
|
|
@ -14,14 +15,13 @@ from lancedb import AsyncConnection, AsyncTable
|
|||
|
||||
from src.config import settings
|
||||
|
||||
from . import QueryResult, VectorRecord, VectorStore
|
||||
from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Schema for LanceDB tables
|
||||
# id: string, vector: fixed_size_list of float32 (1536 dimensions for OpenAI embeddings)
|
||||
# Additional metadata columns are added dynamically
|
||||
VECTOR_DIMENSION = 1536
|
||||
|
||||
# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false, reportUnknownParameterType=false
|
||||
|
||||
|
|
@ -36,17 +36,24 @@ class LanceDBVectorStore(VectorStore):
|
|||
|
||||
_db: AsyncConnection | None = None
|
||||
_db_path: str
|
||||
_db_lock: asyncio.Lock
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the LanceDB vector store."""
|
||||
super().__init__()
|
||||
self._db_path = settings.VECTOR_STORE.LANCEDB_PATH
|
||||
self._db = None
|
||||
self._db_lock = asyncio.Lock()
|
||||
|
||||
async def _get_db(self) -> AsyncConnection:
|
||||
"""Get or create the async database connection."""
|
||||
if self._db is None:
|
||||
self._db = await lancedb.connect_async(self._db_path)
|
||||
"""Get or create the async database connection (asyncio-safe)."""
|
||||
if self._db is not None:
|
||||
return self._db
|
||||
|
||||
async with self._db_lock:
|
||||
# Double-check after acquiring lock
|
||||
if self._db is None:
|
||||
self._db = await lancedb.connect_async(self._db_path)
|
||||
return self._db
|
||||
|
||||
async def _get_table(self, namespace: str) -> AsyncTable | None:
|
||||
|
|
@ -78,7 +85,9 @@ class LanceDBVectorStore(VectorStore):
|
|||
# Create empty table with base schema
|
||||
fields: list[pa.Field] = [
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)),
|
||||
pa.field(
|
||||
"vector", pa.list_(pa.float32(), settings.VECTOR_STORE.DIMENSIONS)
|
||||
),
|
||||
]
|
||||
fields.extend(self._metadata_fields_for_namespace(namespace))
|
||||
schema = pa.schema(fields)
|
||||
|
|
@ -102,7 +111,6 @@ class LanceDBVectorStore(VectorStore):
|
|||
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),
|
||||
]
|
||||
|
||||
if len(parts) == 4:
|
||||
|
|
@ -130,42 +138,11 @@ class LanceDBVectorStore(VectorStore):
|
|||
row[key] = vector.metadata[key]
|
||||
return row
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
namespace: str,
|
||||
vector: VectorRecord,
|
||||
) -> None:
|
||||
"""
|
||||
Upsert a single vector into LanceDB.
|
||||
|
||||
Args:
|
||||
namespace: The namespace (table) to store the vector in
|
||||
vector: VectorRecord containing id, embedding, and metadata
|
||||
"""
|
||||
try:
|
||||
row = self._row_to_dict(vector)
|
||||
table = await self._get_or_create_table(namespace)
|
||||
|
||||
# Use merge_insert for upsert behavior
|
||||
await (
|
||||
table.merge_insert("id")
|
||||
.when_matched_update_all()
|
||||
.when_not_matched_insert_all()
|
||||
.execute([row])
|
||||
)
|
||||
|
||||
logger.debug(f"Upserted vector {vector.id} to namespace {namespace}")
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert vector {vector.id} to namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> None:
|
||||
) -> VectorUpsertResult:
|
||||
"""
|
||||
Upsert multiple vectors into LanceDB.
|
||||
|
||||
|
|
@ -174,7 +151,7 @@ class LanceDBVectorStore(VectorStore):
|
|||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
if not vectors:
|
||||
return
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
|
||||
try:
|
||||
rows = [self._row_to_dict(v) for v in vectors]
|
||||
|
|
@ -189,6 +166,7 @@ class LanceDBVectorStore(VectorStore):
|
|||
)
|
||||
|
||||
logger.debug(f"Upserted {len(vectors)} vectors to namespace {namespace}")
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert {len(vectors)} vectors to namespace {namespace}"
|
||||
|
|
@ -203,7 +181,7 @@ class LanceDBVectorStore(VectorStore):
|
|||
top_k: int = 10,
|
||||
filters: dict[str, Any] | None = None,
|
||||
max_distance: float | None = None,
|
||||
) -> list[QueryResult]:
|
||||
) -> list[VectorQueryResult]:
|
||||
"""
|
||||
Query for similar vectors in LanceDB.
|
||||
|
||||
|
|
@ -215,7 +193,7 @@ class LanceDBVectorStore(VectorStore):
|
|||
max_distance: Optional maximum distance threshold (cosine distance)
|
||||
|
||||
Returns:
|
||||
List of QueryResult objects, ordered by similarity (most similar first)
|
||||
List of VectorQueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
table = await self._get_table(namespace)
|
||||
if table is None:
|
||||
|
|
@ -236,8 +214,8 @@ class LanceDBVectorStore(VectorStore):
|
|||
# LanceDB async API returns list of dicts with incomplete type annotations
|
||||
results = cast(list[dict[str, Any]], await query.to_list())
|
||||
|
||||
# Convert to QueryResult objects
|
||||
query_results: list[QueryResult] = []
|
||||
# Convert to VectorQueryResult objects
|
||||
query_results: list[VectorQueryResult] = []
|
||||
for row in results:
|
||||
dist = float(row.get("_distance", 0.0))
|
||||
|
||||
|
|
@ -253,7 +231,7 @@ class LanceDBVectorStore(VectorStore):
|
|||
}
|
||||
|
||||
query_results.append(
|
||||
QueryResult(
|
||||
VectorQueryResult(
|
||||
id=str(row["id"]),
|
||||
score=dist,
|
||||
metadata=metadata,
|
||||
|
|
@ -343,3 +321,11 @@ class LanceDBVectorStore(VectorStore):
|
|||
except Exception:
|
||||
logger.exception(f"Failed to delete namespace {namespace}")
|
||||
raise
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the LanceDB connection and release resources."""
|
||||
if self._db is not None:
|
||||
# LanceDB AsyncConnection doesn't have an explicit close method,
|
||||
# but we clear the reference to allow garbage collection
|
||||
self._db = None
|
||||
logger.debug("LanceDB connection closed")
|
||||
|
|
|
|||
|
|
@ -0,0 +1,384 @@
|
|||
"""
|
||||
PostgreSQL pgvector vector store implementation.
|
||||
|
||||
This module provides a pgvector-based implementation of the VectorStore interface
|
||||
using the existing embedding columns on documents and message_embeddings tables.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src import models
|
||||
from src.db import SessionLocal
|
||||
|
||||
from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PgVectorStore(VectorStore):
|
||||
"""
|
||||
PostgreSQL pgvector implementation of the VectorStore interface.
|
||||
|
||||
Uses the existing embedding columns on documents and message_embeddings tables,
|
||||
providing transactional consistency with document/message metadata.
|
||||
|
||||
Namespace mapping:
|
||||
- {prefix}.{workspace}.{observer}.{observed} -> documents.embedding
|
||||
- {prefix}.{workspace}.messages -> message_embeddings.embedding
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the pgvector store."""
|
||||
super().__init__()
|
||||
|
||||
def _parse_namespace(self, namespace: str) -> tuple[str, dict[str, str]]:
|
||||
"""
|
||||
Parse a namespace string to determine the table and filter context.
|
||||
|
||||
Args:
|
||||
namespace: Namespace string like "{prefix}.{workspace}.{observer}.{observed}"
|
||||
or "{prefix}.{workspace}.messages"
|
||||
|
||||
Returns:
|
||||
Tuple of (table_type, context_dict) where:
|
||||
- table_type is "documents" or "message_embeddings"
|
||||
- context_dict contains workspace_name and optionally observer/observed
|
||||
"""
|
||||
parts = namespace.split(".")
|
||||
|
||||
# Expected formats:
|
||||
# {prefix}.{workspace}.messages -> message_embeddings
|
||||
# {prefix}.{workspace}.{observer}.{observed} -> documents
|
||||
if len(parts) < 3:
|
||||
raise ValueError(f"Invalid namespace format: {namespace}")
|
||||
|
||||
workspace = parts[1]
|
||||
|
||||
if parts[2] == "messages":
|
||||
return "message_embeddings", {"workspace_name": workspace}
|
||||
elif len(parts) >= 4:
|
||||
observer = parts[2]
|
||||
observed = parts[3]
|
||||
return "documents", {
|
||||
"workspace_name": workspace,
|
||||
"observer": observer,
|
||||
"observed": observed,
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Invalid namespace format: {namespace}")
|
||||
|
||||
async def _get_session(self) -> AsyncSession:
|
||||
"""Get a database session."""
|
||||
return SessionLocal()
|
||||
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> VectorUpsertResult:
|
||||
"""
|
||||
Upsert multiple vectors into the database.
|
||||
|
||||
NOTE: This is a no-op. When pgvector is being used (as primary or secondary),
|
||||
embeddings are written directly to postgres via the ORM (in message.py/document.py).
|
||||
This method exists only to satisfy the VectorStore interface.
|
||||
|
||||
The vector store abstraction is used for:
|
||||
- Queries (which still go through pgvector)
|
||||
- Writing to secondary stores (e.g., turbopuffer during migration)
|
||||
|
||||
Args:
|
||||
namespace: The namespace (determines table)
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
# No-op: embeddings are already written to postgres via ORM
|
||||
if vectors:
|
||||
logger.debug(
|
||||
f"PgVectorStore.upsert_many() no-op for {len(vectors)} vectors in {namespace} (embeddings written via ORM)"
|
||||
)
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
|
||||
async def query(
|
||||
self,
|
||||
namespace: str,
|
||||
embedding: list[float],
|
||||
*,
|
||||
top_k: int = 10,
|
||||
filters: dict[str, Any] | None = None,
|
||||
max_distance: float | None = None,
|
||||
) -> list[VectorQueryResult]:
|
||||
"""
|
||||
Query for similar vectors using pgvector cosine distance.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to query
|
||||
embedding: The query embedding vector
|
||||
top_k: Maximum number of results to return
|
||||
filters: Optional metadata filters
|
||||
max_distance: Optional maximum distance threshold (cosine distance)
|
||||
|
||||
Returns:
|
||||
List of VectorQueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
table_type, context = self._parse_namespace(namespace)
|
||||
|
||||
db = await self._get_session()
|
||||
try:
|
||||
if table_type == "documents":
|
||||
results = await self._query_documents(
|
||||
db, context, embedding, top_k, filters, max_distance
|
||||
)
|
||||
elif table_type == "message_embeddings":
|
||||
results = await self._query_message_embeddings(
|
||||
db, context, embedding, top_k, filters, max_distance
|
||||
)
|
||||
else:
|
||||
results = []
|
||||
|
||||
logger.debug(
|
||||
f"Query returned {len(results)} results from namespace {namespace}"
|
||||
)
|
||||
return results
|
||||
|
||||
except Exception:
|
||||
logger.exception(f"Failed to query namespace {namespace}")
|
||||
raise
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
async def _query_documents(
|
||||
self,
|
||||
db: AsyncSession,
|
||||
context: dict[str, str],
|
||||
embedding: list[float],
|
||||
top_k: int,
|
||||
filters: dict[str, Any] | None,
|
||||
max_distance: float | None,
|
||||
) -> list[VectorQueryResult]:
|
||||
"""Query documents table for similar vectors."""
|
||||
# Build the query with cosine distance
|
||||
# pgvector uses <=> for cosine distance
|
||||
stmt = (
|
||||
select(
|
||||
models.Document.id,
|
||||
models.Document.embedding.cosine_distance(embedding).label("distance"),
|
||||
models.Document.workspace_name,
|
||||
models.Document.observer,
|
||||
models.Document.observed,
|
||||
models.Document.session_name,
|
||||
models.Document.level,
|
||||
)
|
||||
.where(models.Document.embedding.isnot(None))
|
||||
.where(models.Document.deleted_at.is_(None))
|
||||
.where(models.Document.workspace_name == context["workspace_name"])
|
||||
.where(models.Document.observer == context["observer"])
|
||||
.where(models.Document.observed == context["observed"])
|
||||
)
|
||||
|
||||
# Apply additional filters
|
||||
if filters:
|
||||
if "session_name" in filters:
|
||||
stmt = stmt.where(
|
||||
models.Document.session_name == filters["session_name"]
|
||||
)
|
||||
if "level" in filters:
|
||||
stmt = stmt.where(models.Document.level == filters["level"])
|
||||
|
||||
# Apply max_distance filter
|
||||
if max_distance is not None:
|
||||
stmt = stmt.where(
|
||||
models.Document.embedding.cosine_distance(embedding) <= max_distance
|
||||
)
|
||||
|
||||
# Order by distance and limit
|
||||
stmt = stmt.order_by("distance").limit(top_k)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = result.all()
|
||||
|
||||
return [
|
||||
VectorQueryResult(
|
||||
id=str(row.id),
|
||||
score=float(row.distance),
|
||||
metadata={
|
||||
"workspace_name": row.workspace_name,
|
||||
"observer": row.observer,
|
||||
"observed": row.observed,
|
||||
"session_name": row.session_name,
|
||||
"level": row.level,
|
||||
},
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def _query_message_embeddings(
|
||||
self,
|
||||
db: AsyncSession,
|
||||
context: dict[str, str],
|
||||
embedding: list[float],
|
||||
top_k: int,
|
||||
filters: dict[str, Any] | None,
|
||||
max_distance: float | None,
|
||||
) -> list[VectorQueryResult]:
|
||||
"""Query message_embeddings table for similar vectors."""
|
||||
# Build the query with cosine distance
|
||||
stmt = (
|
||||
select(
|
||||
models.MessageEmbedding.id,
|
||||
models.MessageEmbedding.embedding.cosine_distance(embedding).label(
|
||||
"distance"
|
||||
),
|
||||
models.MessageEmbedding.message_id,
|
||||
models.MessageEmbedding.workspace_name,
|
||||
models.MessageEmbedding.session_name,
|
||||
models.MessageEmbedding.peer_name,
|
||||
)
|
||||
.where(models.MessageEmbedding.embedding.isnot(None))
|
||||
.where(models.MessageEmbedding.workspace_name == context["workspace_name"])
|
||||
)
|
||||
|
||||
# Apply additional filters
|
||||
if filters:
|
||||
if "session_name" in filters:
|
||||
stmt = stmt.where(
|
||||
models.MessageEmbedding.session_name == filters["session_name"]
|
||||
)
|
||||
if "peer_name" in filters:
|
||||
stmt = stmt.where(
|
||||
models.MessageEmbedding.peer_name == filters["peer_name"]
|
||||
)
|
||||
if "message_id" in filters:
|
||||
stmt = stmt.where(
|
||||
models.MessageEmbedding.message_id == filters["message_id"]
|
||||
)
|
||||
|
||||
# Apply max_distance filter
|
||||
if max_distance is not None:
|
||||
stmt = stmt.where(
|
||||
models.MessageEmbedding.embedding.cosine_distance(embedding)
|
||||
<= max_distance
|
||||
)
|
||||
|
||||
# Order by distance and limit
|
||||
stmt = stmt.order_by("distance").limit(top_k)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = result.all()
|
||||
|
||||
return [
|
||||
VectorQueryResult(
|
||||
id=str(row.id),
|
||||
score=float(row.distance),
|
||||
metadata={
|
||||
"embedding_id": row.id,
|
||||
"message_id": row.message_id,
|
||||
"workspace_name": row.workspace_name,
|
||||
"session_name": row.session_name,
|
||||
"peer_name": row.peer_name,
|
||||
},
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def delete_many(self, namespace: str, ids: list[str]) -> None:
|
||||
"""
|
||||
Delete vectors by setting embedding to NULL.
|
||||
|
||||
Args:
|
||||
namespace: The namespace containing the vectors
|
||||
ids: List of vector identifiers to delete
|
||||
"""
|
||||
if not ids:
|
||||
return
|
||||
|
||||
table_type, _ = self._parse_namespace(namespace)
|
||||
|
||||
db = await self._get_session()
|
||||
try:
|
||||
if table_type == "documents":
|
||||
stmt = (
|
||||
update(models.Document)
|
||||
.where(models.Document.id.in_(ids))
|
||||
.values(embedding=None)
|
||||
)
|
||||
await db.execute(stmt)
|
||||
|
||||
elif table_type == "message_embeddings":
|
||||
for vector_id in ids:
|
||||
try:
|
||||
embedding_id = int(vector_id)
|
||||
except ValueError as exc:
|
||||
raise ValueError(
|
||||
f"Invalid message vector id format: {vector_id}"
|
||||
) from exc
|
||||
|
||||
stmt = (
|
||||
update(models.MessageEmbedding)
|
||||
.where(models.MessageEmbedding.id == embedding_id)
|
||||
.values(embedding=None)
|
||||
)
|
||||
await db.execute(stmt)
|
||||
|
||||
await db.commit()
|
||||
logger.debug(
|
||||
f"Deleted {len(ids)} vectors from {table_type} in namespace {namespace}"
|
||||
)
|
||||
|
||||
except Exception:
|
||||
await db.rollback()
|
||||
logger.exception(
|
||||
f"Failed to delete {len(ids)} vectors from namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
async def delete_namespace(self, namespace: str) -> None:
|
||||
"""
|
||||
Delete all vectors in a namespace by setting embedding to NULL.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to delete
|
||||
"""
|
||||
table_type, context = self._parse_namespace(namespace)
|
||||
|
||||
db = await self._get_session()
|
||||
try:
|
||||
if table_type == "documents":
|
||||
stmt = (
|
||||
update(models.Document)
|
||||
.where(models.Document.workspace_name == context["workspace_name"])
|
||||
.where(models.Document.observer == context["observer"])
|
||||
.where(models.Document.observed == context["observed"])
|
||||
.values(embedding=None)
|
||||
)
|
||||
await db.execute(stmt)
|
||||
|
||||
elif table_type == "message_embeddings":
|
||||
stmt = (
|
||||
update(models.MessageEmbedding)
|
||||
.where(
|
||||
models.MessageEmbedding.workspace_name
|
||||
== context["workspace_name"]
|
||||
)
|
||||
.values(embedding=None)
|
||||
)
|
||||
await db.execute(stmt)
|
||||
|
||||
await db.commit()
|
||||
logger.debug(f"Deleted all vectors from namespace {namespace}")
|
||||
|
||||
except Exception:
|
||||
await db.rollback()
|
||||
logger.exception(f"Failed to delete namespace {namespace}")
|
||||
raise
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the pgvector store (no-op for pgvector)"""
|
||||
pass
|
||||
|
|
@ -15,7 +15,7 @@ from turbopuffer.types import Filter
|
|||
|
||||
from src.config import settings
|
||||
|
||||
from . import QueryResult, VectorRecord, VectorStore
|
||||
from . import VectorQueryResult, VectorRecord, VectorStore, VectorUpsertResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -60,44 +60,11 @@ class TurbopufferVectorStore(VectorStore):
|
|||
"""Get a Turbopuffer namespace object."""
|
||||
return self.tpuf.namespace(namespace)
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
namespace: str,
|
||||
vector: VectorRecord,
|
||||
) -> None:
|
||||
"""
|
||||
Upsert a single vector into Turbopuffer.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vector in
|
||||
vector: VectorRecord containing id, embedding, and optional metadata
|
||||
"""
|
||||
ns = self._get_namespace(namespace)
|
||||
attributes = vector.metadata or {}
|
||||
|
||||
try:
|
||||
# Build row data
|
||||
row: dict[str, Any] = {
|
||||
"id": vector.id,
|
||||
"vector": vector.embedding,
|
||||
**attributes,
|
||||
}
|
||||
|
||||
await ns.write(
|
||||
upsert_rows=[row],
|
||||
distance_metric=DISTANCE_METRIC,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert vector {vector.id} to namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> None:
|
||||
) -> VectorUpsertResult:
|
||||
"""
|
||||
Upsert multiple vectors into Turbopuffer.
|
||||
|
||||
|
|
@ -106,7 +73,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
if not vectors:
|
||||
return
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
|
|
@ -124,6 +91,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
upsert_rows=rows,
|
||||
distance_metric=DISTANCE_METRIC,
|
||||
)
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert {len(vectors)} vectors to namespace {namespace}"
|
||||
|
|
@ -138,7 +106,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
top_k: int = 10,
|
||||
filters: dict[str, Any] | None = None,
|
||||
max_distance: float | None = None,
|
||||
) -> list[QueryResult]:
|
||||
) -> list[VectorQueryResult]:
|
||||
"""
|
||||
Query for similar vectors in Turbopuffer.
|
||||
|
||||
|
|
@ -150,7 +118,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
max_distance: Optional maximum distance threshold (cosine distance)
|
||||
|
||||
Returns:
|
||||
List of QueryResult objects, ordered by similarity (most similar first)
|
||||
List of VectorQueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
|
|
@ -178,7 +146,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
|
||||
response = await ns.query(**query_kwargs)
|
||||
|
||||
query_results: list[QueryResult] = []
|
||||
query_results: list[VectorQueryResult] = []
|
||||
for row in response.rows or []:
|
||||
# Distance is accessed via row["$dist"]
|
||||
dist: float = float(row["$dist"]) if "$dist" in row else 0.0
|
||||
|
|
@ -197,7 +165,7 @@ class TurbopufferVectorStore(VectorStore):
|
|||
}
|
||||
|
||||
query_results.append(
|
||||
QueryResult(
|
||||
VectorQueryResult(
|
||||
id=str(row.id),
|
||||
score=dist,
|
||||
metadata=row_metadata,
|
||||
|
|
@ -295,3 +263,8 @@ class TurbopufferVectorStore(VectorStore):
|
|||
except Exception:
|
||||
logger.exception(f"Failed to delete namespace {namespace}")
|
||||
raise
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the Turbopuffer client and release resources."""
|
||||
await self.tpuf.close()
|
||||
logger.debug("Turbopuffer client closed")
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Hooks for revision f1a2b3c4d5e6 (add_chunk_index_to_message_embeddings)."""
|
||||
"""Hooks for revision f1a2b3c4d5e6 (support_external_embeddings)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -7,20 +7,16 @@ from tests.alembic.verifier import MigrationVerifier
|
|||
|
||||
|
||||
@register_before_upgrade("f1a2b3c4d5e6")
|
||||
def prepare_add_chunk_index_to_message_embeddings(
|
||||
def prepare_support_external_embeddings(
|
||||
verifier: MigrationVerifier,
|
||||
) -> None:
|
||||
"""Seed state and assertions before upgrading to f1a2b3c4d5e6."""
|
||||
verifier.assert_column_exists("message_embeddings", "embedding", nullable=False)
|
||||
# Verify chunk_index column doesn't exist before migration
|
||||
verifier.assert_column_exists("message_embeddings", "chunk_index", exists=False)
|
||||
|
||||
|
||||
@register_after_upgrade("f1a2b3c4d5e6")
|
||||
def verify_add_chunk_index_to_message_embeddings(
|
||||
def verify_support_external_embeddings(
|
||||
verifier: MigrationVerifier,
|
||||
) -> None:
|
||||
"""Add assertions validating the effects of f1a2b3c4d5e6."""
|
||||
# Verify chunk_index column was added with correct properties
|
||||
verifier.assert_column_exists("message_embeddings", "chunk_index", nullable=False)
|
||||
verifier.assert_column_exists("message_embeddings", "embedding", nullable=True)
|
||||
|
|
|
|||
|
|
@ -381,34 +381,32 @@ def mock_vector_store():
|
|||
"""Mock vector store operations for testing"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from src.vector_store import QueryResult, VectorRecord
|
||||
from src.vector_store import VectorQueryResult, VectorRecord, VectorUpsertResult
|
||||
|
||||
# Create a mock vector store that stores vectors in memory
|
||||
vector_storage: dict[str, dict[str, tuple[list[float], dict[str, Any]]]] = {}
|
||||
|
||||
async def mock_upsert(namespace: str, vector: VectorRecord) -> None:
|
||||
if namespace not in vector_storage:
|
||||
vector_storage[namespace] = {}
|
||||
vector_storage[namespace][vector.id] = (vector.embedding, vector.metadata)
|
||||
|
||||
async def mock_upsert_many(namespace: str, vectors: list[VectorRecord]) -> None:
|
||||
async def mock_upsert_many(
|
||||
namespace: str, vectors: list[VectorRecord]
|
||||
) -> VectorUpsertResult:
|
||||
if namespace not in vector_storage:
|
||||
vector_storage[namespace] = {}
|
||||
for vector in vectors:
|
||||
vector_storage[namespace][vector.id] = (vector.embedding, vector.metadata)
|
||||
return VectorUpsertResult(primary_ok=True)
|
||||
|
||||
async def mock_query(
|
||||
namespace: str, embedding: list[float], **kwargs: Any
|
||||
) -> list[QueryResult]:
|
||||
) -> list[VectorQueryResult]:
|
||||
_ = embedding # unused in mock
|
||||
if namespace not in vector_storage:
|
||||
return []
|
||||
|
||||
# Simple mock: return all vectors in the namespace as results
|
||||
results: list[QueryResult] = []
|
||||
results: list[VectorQueryResult] = []
|
||||
for vec_id, (_vec_embedding, metadata) in vector_storage[namespace].items():
|
||||
results.append(
|
||||
QueryResult(
|
||||
VectorQueryResult(
|
||||
id=vec_id,
|
||||
score=0.1, # Mock score
|
||||
metadata=metadata,
|
||||
|
|
@ -425,24 +423,51 @@ def mock_vector_store():
|
|||
async def mock_delete_namespace(namespace: str) -> None:
|
||||
vector_storage.pop(namespace, None)
|
||||
|
||||
# Clear the cache on get_vector_store before patching
|
||||
from src.vector_store import get_vector_store
|
||||
|
||||
get_vector_store.cache_clear() # type: ignore
|
||||
|
||||
# Create the mock vector store
|
||||
mock_vs = MagicMock()
|
||||
mock_vs.upsert_many = AsyncMock(side_effect=mock_upsert_many)
|
||||
mock_vs.query = AsyncMock(side_effect=mock_query)
|
||||
mock_vs.delete_many = AsyncMock(side_effect=mock_delete_many)
|
||||
mock_vs.delete_namespace = AsyncMock(side_effect=mock_delete_namespace)
|
||||
|
||||
def mock_get_vector_namespace(
|
||||
namespace_type: str,
|
||||
workspace_name: str,
|
||||
observer: str | None = None,
|
||||
observed: str | None = None,
|
||||
) -> str:
|
||||
if namespace_type == "document":
|
||||
if observer is None or observed is None:
|
||||
raise ValueError(
|
||||
"observer and observed are required for document namespaces"
|
||||
)
|
||||
return f"honcho2345.{workspace_name}.{observer}.{observed}"
|
||||
if namespace_type == "message":
|
||||
return f"honcho2345.{workspace_name}.messages"
|
||||
raise ValueError(f"Unknown namespace type: {namespace_type}")
|
||||
|
||||
mock_vs.get_vector_namespace = mock_get_vector_namespace
|
||||
|
||||
with (
|
||||
patch("src.vector_store.get_vector_store") as mock_get_vs,
|
||||
patch("src.crud.document.get_vector_store", return_value=mock_vs),
|
||||
patch("src.crud.workspace.get_vector_store", return_value=mock_vs),
|
||||
patch("src.crud.session.get_vector_store", return_value=mock_vs),
|
||||
patch("src.crud.message.get_vector_store", return_value=mock_vs),
|
||||
patch(
|
||||
"src.deriver.vector_reconciliation.get_vector_store", return_value=mock_vs
|
||||
),
|
||||
patch("src.utils.search.get_vector_store", return_value=mock_vs),
|
||||
):
|
||||
mock_vs = MagicMock()
|
||||
mock_vs.upsert = AsyncMock(side_effect=mock_upsert)
|
||||
mock_vs.upsert_many = AsyncMock(side_effect=mock_upsert_many)
|
||||
mock_vs.query = AsyncMock(side_effect=mock_query)
|
||||
mock_vs.delete_many = AsyncMock(side_effect=mock_delete_many)
|
||||
mock_vs.delete_namespace = AsyncMock(side_effect=mock_delete_namespace)
|
||||
mock_vs.get_document_namespace = (
|
||||
lambda ws, obs, obd: f"honcho:{ws}:{obs}:{obd}" # pyright: ignore[reportUnknownLambdaType]
|
||||
)
|
||||
mock_vs.get_message_namespace = lambda ws: f"honcho:{ws}:messages" # pyright: ignore[reportUnknownLambdaType]
|
||||
|
||||
mock_get_vs.return_value = mock_vs
|
||||
|
||||
yield mock_vs
|
||||
|
||||
# Clear cache after test as well for cleanliness
|
||||
get_vector_store.cache_clear() # type: ignore
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_llm_call_functions():
|
||||
|
|
|
|||
|
|
@ -329,5 +329,3 @@ async def test_message_chunking_creates_multiple_embeddings(
|
|||
assert embedding_record.workspace_name == test_workspace.name
|
||||
assert embedding_record.session_name == test_session.name
|
||||
assert embedding_record.peer_name == test_peer.name
|
||||
# chunk_index should be set for each chunk
|
||||
assert embedding_record.chunk_index is not None
|
||||
|
|
|
|||
|
|
@ -0,0 +1,51 @@
|
|||
"""Tests for vector reconciliation configuration logic.
|
||||
|
||||
These are unit tests that test configuration logic without requiring database
|
||||
or vector store fixtures.
|
||||
"""
|
||||
|
||||
from src.config import VectorStoreSettings
|
||||
|
||||
|
||||
def test_reconciliation_enabled_pgvector_primary():
|
||||
"""Reconciliation enabled when pgvector is primary with secondary"""
|
||||
settings = VectorStoreSettings(
|
||||
PRIMARY_TYPE="pgvector",
|
||||
SECONDARY_TYPE="turbopuffer",
|
||||
TURBOPUFFER_API_KEY="test-key",
|
||||
)
|
||||
assert settings.should_run_reconciliation is True
|
||||
|
||||
|
||||
def test_reconciliation_enabled_pgvector_secondary():
|
||||
"""Reconciliation enabled when pgvector is secondary"""
|
||||
settings = VectorStoreSettings(
|
||||
PRIMARY_TYPE="turbopuffer",
|
||||
SECONDARY_TYPE="pgvector",
|
||||
TURBOPUFFER_API_KEY="test-key",
|
||||
)
|
||||
assert settings.should_run_reconciliation is True
|
||||
|
||||
|
||||
def test_reconciliation_disabled_single_store():
|
||||
"""Reconciliation disabled when no secondary configured"""
|
||||
settings = VectorStoreSettings(
|
||||
PRIMARY_TYPE="turbopuffer", TURBOPUFFER_API_KEY="test-key", SECONDARY_TYPE=None
|
||||
)
|
||||
assert settings.should_run_reconciliation is False
|
||||
|
||||
|
||||
def test_reconciliation_disabled_no_pgvector():
|
||||
"""Reconciliation disabled when both stores are non-pgvector"""
|
||||
settings = VectorStoreSettings(
|
||||
PRIMARY_TYPE="turbopuffer",
|
||||
SECONDARY_TYPE="lancedb",
|
||||
TURBOPUFFER_API_KEY="test-key",
|
||||
)
|
||||
assert settings.should_run_reconciliation is False
|
||||
|
||||
|
||||
def test_reconciliation_disabled_pgvector_only():
|
||||
"""Reconciliation disabled when pgvector is primary but no secondary"""
|
||||
settings = VectorStoreSettings(PRIMARY_TYPE="pgvector", SECONDARY_TYPE=None)
|
||||
assert settings.should_run_reconciliation is False
|
||||
Loading…
Reference in New Issue