feat: init turbopuffer and lanceDB
This commit is contained in:
parent
4ee2f8bd0c
commit
ca218da615
|
|
@ -0,0 +1,136 @@
|
|||
"""remove embedding columns for vector store migration
|
||||
|
||||
This migration removes the embedding columns from the message_embeddings and documents
|
||||
tables as part of the migration from pgvector to external vector stores (turbopuffer/lancedb).
|
||||
|
||||
Revision ID: f1a2b3c4d5e6
|
||||
Revises: baa22cad81e2
|
||||
Create Date: 2025-11-24 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from pgvector.sqlalchemy import Vector
|
||||
|
||||
from migrations.utils import column_exists, get_schema, index_exists
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f1a2b3c4d5e6"
|
||||
down_revision: str | None = "baa22cad81e2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
schema = get_schema()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Remove embedding columns and HNSW indexes from message_embeddings and documents tables.
|
||||
Also add chunk_index column to message_embeddings for tracking chunked message embeddings.
|
||||
"""
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
|
||||
# === message_embeddings table ===
|
||||
|
||||
# Drop HNSW index on message_embeddings.embedding if it exists
|
||||
# Check for both possible index names (old naming vs new naming convention)
|
||||
for index_name in [
|
||||
"ix_message_embeddings_embedding_hnsw",
|
||||
"idx_message_embeddings_embedding_hnsw",
|
||||
]:
|
||||
if index_exists("message_embeddings", index_name, inspector):
|
||||
op.drop_index(index_name, table_name="message_embeddings", schema=schema)
|
||||
|
||||
# Drop embedding column from message_embeddings if it exists
|
||||
if column_exists("message_embeddings", "embedding", inspector):
|
||||
op.drop_column("message_embeddings", "embedding", schema=schema)
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
# === documents table ===
|
||||
|
||||
# Drop HNSW index on documents.embedding if it exists
|
||||
# Check for both possible index names (old naming vs new naming convention)
|
||||
for index_name in [
|
||||
"ix_documents_embedding_hnsw",
|
||||
"idx_documents_embedding_hnsw",
|
||||
]:
|
||||
if index_exists("documents", index_name, inspector):
|
||||
op.drop_index(index_name, table_name="documents", schema=schema)
|
||||
|
||||
# Drop embedding column from documents if it exists
|
||||
if column_exists("documents", "embedding", inspector):
|
||||
op.drop_column("documents", "embedding", schema=schema)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Restore embedding columns and HNSW indexes, remove chunk_index.
|
||||
|
||||
Note: This downgrade will create empty embedding columns. The actual embeddings
|
||||
would need to be restored from a backup or regenerated if rolling back this migration.
|
||||
"""
|
||||
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
|
||||
# === documents table ===
|
||||
|
||||
# Add embedding column back to documents if it doesn't exist
|
||||
if not column_exists("documents", "embedding", inspector):
|
||||
op.add_column(
|
||||
"documents",
|
||||
sa.Column("embedding", Vector(1536), nullable=True),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Recreate HNSW index on documents.embedding
|
||||
if not index_exists("documents", "ix_documents_embedding_hnsw", inspector):
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE INDEX ix_documents_embedding_hnsw
|
||||
ON {schema}.documents
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
WITH (m = 16, ef_construction = 64)
|
||||
"""
|
||||
)
|
||||
|
||||
# === message_embeddings table ===
|
||||
|
||||
# Remove chunk_index column if it exists
|
||||
if column_exists("message_embeddings", "chunk_index", inspector):
|
||||
op.drop_column("message_embeddings", "chunk_index", schema=schema)
|
||||
|
||||
# Add embedding column back to message_embeddings if it doesn't exist
|
||||
if not column_exists("message_embeddings", "embedding", inspector):
|
||||
op.add_column(
|
||||
"message_embeddings",
|
||||
sa.Column("embedding", Vector(1536), nullable=True),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
# Recreate HNSW index on message_embeddings.embedding
|
||||
if not index_exists(
|
||||
"message_embeddings", "ix_message_embeddings_embedding_hnsw", inspector
|
||||
):
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE INDEX ix_message_embeddings_embedding_hnsw
|
||||
ON {schema}.message_embeddings
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
WITH (m = 16, ef_construction = 64)
|
||||
"""
|
||||
)
|
||||
|
|
@ -35,6 +35,8 @@ dependencies = [
|
|||
"json-repair>=0.49.0",
|
||||
"redis>=6.0.0",
|
||||
"cashews[redis]==7.4.1",
|
||||
"turbopuffer>=1.8.1",
|
||||
"lancedb>=0.25.3",
|
||||
]
|
||||
[tool.uv]
|
||||
dev-dependencies = [
|
||||
|
|
@ -110,7 +112,7 @@ reportUnusedCallResult = false
|
|||
reportCallInDefaultInitializer = false
|
||||
reportAny = false
|
||||
reportExplicitAny = false
|
||||
allowedUntypedLibraries = ["langfuse"]
|
||||
allowedUntypedLibraries = ["langfuse", "lancedb", "pyarrow"]
|
||||
reportImplicitOverride = false
|
||||
reportImportCycles = false
|
||||
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ class TomlConfigSettingsSource(PydanticBaseSettingsSource):
|
|||
"SUMMARY": "summary",
|
||||
"WEBHOOK": "webhook",
|
||||
"DREAM": "dream",
|
||||
"VECTOR_STORE": "vector_store",
|
||||
"": "app", # For AppSettings with no prefix
|
||||
}
|
||||
|
||||
|
|
@ -356,6 +357,38 @@ class DreamSettings(BackupLLMSettingsMixin, HonchoSettings):
|
|||
MAX_OUTPUT_TOKENS: Annotated[int, Field(default=2000, gt=0, le=10_000)] = 2000
|
||||
|
||||
|
||||
class VectorStoreSettings(HonchoSettings):
|
||||
"""Settings for external vector store (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"
|
||||
|
||||
# Global namespace prefix for all vector namespaces
|
||||
# Namespaces follow the pattern:
|
||||
# - Documents: {NAMESPACE}-{workspace}-{observer}-{observed}
|
||||
# - Messages: {NAMESPACE}-{workspace}-messages
|
||||
NAMESPACE: str = "honcho"
|
||||
|
||||
# Turbopuffer-specific settings
|
||||
TURBOPUFFER_API_KEY: str | None = None
|
||||
# Turbopuffer region (e.g., "gcp-us-east4", "aws-us-east-1")
|
||||
# Can also be set via TURBOPUFFER_REGION environment variable
|
||||
TURBOPUFFER_REGION: str | None = None
|
||||
|
||||
# LanceDB-specific settings (local embedded mode)
|
||||
LANCEDB_PATH: str = "./lancedb_data"
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _require_api_key_for_turbopuffer(self) -> "VectorStoreSettings":
|
||||
if self.TYPE == "turbopuffer" and not self.TURBOPUFFER_API_KEY:
|
||||
raise ValueError(
|
||||
"VECTOR_STORE_TURBOPUFFER_API_KEY must be set when TYPE is 'turbopuffer'"
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
class AppSettings(HonchoSettings):
|
||||
# No env_prefix for app-level settings
|
||||
model_config = SettingsConfigDict( # pyright: ignore
|
||||
|
|
@ -397,6 +430,7 @@ class AppSettings(HonchoSettings):
|
|||
METRICS: MetricsSettings = Field(default_factory=MetricsSettings)
|
||||
CACHE: CacheSettings = Field(default_factory=CacheSettings)
|
||||
DREAM: DreamSettings = Field(default_factory=DreamSettings)
|
||||
VECTOR_STORE: VectorStoreSettings = Field(default_factory=VectorStoreSettings)
|
||||
|
||||
@field_validator("LOG_LEVEL")
|
||||
def validate_log_level(cls, v: str) -> str:
|
||||
|
|
@ -409,13 +443,16 @@ class AppSettings(HonchoSettings):
|
|||
def propagate_namespace(self) -> "AppSettings":
|
||||
"""Propagate top-level NAMESPACE to nested settings if not explicitly set.
|
||||
|
||||
After this validator runs, CACHE.NAMESPACE and METRICS.NAMESPACE are guaranteed
|
||||
to exist.
|
||||
After this validator runs, CACHE.NAMESPACE, METRICS.NAMESPACE, and
|
||||
VECTOR_STORE.NAMESPACE are guaranteed to exist.
|
||||
"""
|
||||
if self.CACHE.NAMESPACE is None:
|
||||
self.CACHE.NAMESPACE = self.NAMESPACE
|
||||
if self.METRICS.NAMESPACE is None:
|
||||
self.METRICS.NAMESPACE = self.NAMESPACE
|
||||
# Note: VECTOR_STORE.NAMESPACE has its own default of "honcho",
|
||||
# but we propagate the top-level NAMESPACE if the user explicitly set it
|
||||
# and wants consistency across all namespaced services
|
||||
return self
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from src import models, schemas
|
|||
from src.config import settings
|
||||
from src.embedding_client import embedding_client
|
||||
from src.exceptions import ValidationException
|
||||
from src.utils.filter import apply_filter
|
||||
from src.vector_store import VectorRecord, get_vector_store
|
||||
|
||||
logger = getLogger(__name__)
|
||||
|
||||
|
|
@ -61,7 +61,7 @@ async def query_documents(
|
|||
query: Search query text
|
||||
observer: Name of the observing peer
|
||||
observed: Name of the observed peer
|
||||
filters: Optional filters to apply
|
||||
filters: Optional filters to apply at vector store level (supports: level, session_name)
|
||||
max_distance: Maximum cosine distance for results
|
||||
top_k: Number of results to return
|
||||
embedding: Optional pre-computed embedding for the query (avoids extra API call if possible)
|
||||
|
|
@ -78,22 +78,56 @@ async def query_documents(
|
|||
f"Query exceeds maximum token limit of {settings.MAX_EMBEDDING_TOKENS}."
|
||||
) from e
|
||||
|
||||
# Get vector store and namespace for this collection
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(workspace_name, observer, observed)
|
||||
|
||||
# Build vector store filters
|
||||
# Convert filter dict to vector store format (handles level, session_name, etc.)
|
||||
vector_filters: dict[str, Any] = {}
|
||||
if filters:
|
||||
# Direct pass-through for simple equality filters
|
||||
# The filters dict can contain: level, session_name, or other document fields
|
||||
# We can push level and session_name to vector store since they're in metadata
|
||||
for key in ["level", "session_name"]:
|
||||
if key in filters:
|
||||
vector_filters[key] = filters[key]
|
||||
|
||||
# Query vector store for similar documents with filters applied
|
||||
vector_results = await vector_store.query(
|
||||
namespace,
|
||||
embedding,
|
||||
top_k=top_k,
|
||||
max_distance=max_distance,
|
||||
filters=vector_filters if vector_filters else None,
|
||||
)
|
||||
|
||||
if not vector_results:
|
||||
return []
|
||||
|
||||
# Get document IDs from vector results (vector ID = document ID for documents)
|
||||
document_ids = [result.id for result in vector_results]
|
||||
|
||||
# Fetch documents from database
|
||||
# No additional filtering needed since vector store already applied all supported filters
|
||||
stmt = (
|
||||
select(models.Document)
|
||||
.where(models.Document.workspace_name == workspace_name)
|
||||
.where(models.Document.observer == observer)
|
||||
.where(models.Document.observed == observed)
|
||||
.where(models.Document.id.in_(document_ids))
|
||||
)
|
||||
if max_distance is not None:
|
||||
stmt = stmt.where(
|
||||
models.Document.embedding.cosine_distance(embedding) < max_distance
|
||||
)
|
||||
stmt = apply_filter(stmt, models.Document, filters)
|
||||
stmt = stmt.limit(top_k).order_by(
|
||||
models.Document.embedding.cosine_distance(embedding)
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
return result.scalars().all()
|
||||
documents = {doc.id: doc for doc in result.scalars().all()}
|
||||
|
||||
# Return documents in order of similarity (preserving vector store order)
|
||||
ordered_docs: list[models.Document] = []
|
||||
for vr in vector_results:
|
||||
if vr.id in documents:
|
||||
ordered_docs.append(documents[vr.id])
|
||||
|
||||
return ordered_docs
|
||||
|
||||
|
||||
async def create_documents(
|
||||
|
|
@ -119,6 +153,8 @@ async def create_documents(
|
|||
Count of new documents
|
||||
"""
|
||||
honcho_documents: list[models.Document] = []
|
||||
embeddings_to_store: list[tuple[str, list[float]]] = [] # [(doc_id, embedding)]
|
||||
|
||||
for doc in documents:
|
||||
try:
|
||||
# for each document, if deduplicate is True, perform a process
|
||||
|
|
@ -132,27 +168,59 @@ async def create_documents(
|
|||
continue
|
||||
|
||||
metadata_dict = doc.metadata.model_dump(exclude_none=True)
|
||||
honcho_documents.append(
|
||||
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,
|
||||
embedding=doc.embedding,
|
||||
session_name=doc.session_name,
|
||||
)
|
||||
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,
|
||||
)
|
||||
honcho_documents.append(new_doc)
|
||||
|
||||
# Track embedding for vector store (will use document's generated ID)
|
||||
if doc.embedding:
|
||||
embeddings_to_store.append((new_doc.id, doc.embedding))
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error adding new document to {workspace_name}/{doc.session_name}/{observer}/{observed}: {e}"
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
db.add_all(honcho_documents)
|
||||
await db.commit()
|
||||
|
||||
# Store embeddings in vector store after documents are committed
|
||||
if embeddings_to_store:
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, observer, observed
|
||||
)
|
||||
|
||||
# Build vector records with metadata for filtering
|
||||
vector_records: list[VectorRecord] = []
|
||||
doc_lookup = {doc.id: doc for doc in honcho_documents}
|
||||
for doc_id, embedding in embeddings_to_store:
|
||||
doc = doc_lookup[doc_id]
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=doc_id,
|
||||
embedding=embedding,
|
||||
metadata={
|
||||
"workspace_name": workspace_name,
|
||||
"observer": observer,
|
||||
"observed": observed,
|
||||
"session_name": doc.session_name,
|
||||
"level": doc.level,
|
||||
},
|
||||
)
|
||||
)
|
||||
await vector_store.upsert_many(namespace, vector_records)
|
||||
|
||||
except IntegrityError as e:
|
||||
await db.rollback()
|
||||
raise ValidationException(
|
||||
|
|
@ -216,8 +284,17 @@ async def is_rejected_duplicate(
|
|||
logger.warning(
|
||||
f"[DUPLICATE DETECTION] Deleting existing in favor of new. new='{doc.content}', existing='{existing_doc.content}'."
|
||||
)
|
||||
# Delete from database
|
||||
await db.delete(existing_doc)
|
||||
await db.flush() # Flush to make deletion visible in this transaction
|
||||
|
||||
# Delete from vector store
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_document_namespace(
|
||||
workspace_name, observer, observed
|
||||
)
|
||||
await vector_store.delete(namespace, existing_doc.id)
|
||||
|
||||
return False # Don't reject the new document
|
||||
|
||||
# Existing document has more information, reject the new one
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from src import models, schemas
|
|||
from src.config import settings
|
||||
from src.embedding_client import embedding_client
|
||||
from src.utils.filter import apply_filter
|
||||
from src.vector_store import VectorRecord, get_vector_store
|
||||
|
||||
from .session import get_or_create_session
|
||||
|
||||
|
|
@ -136,25 +137,52 @@ async def create_messages(
|
|||
}
|
||||
embedding_dict = await embedding_client.batch_embed(id_resource_dict)
|
||||
|
||||
# Create MessageEmbedding entries for each embedded message
|
||||
# Get vector store and namespace for this workspace's messages
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_message_namespace(workspace_name)
|
||||
|
||||
# Create MessageEmbedding entries and vector records
|
||||
embedding_objects: list[models.MessageEmbedding] = []
|
||||
vector_records: list[VectorRecord] = []
|
||||
|
||||
for message_obj in message_objects:
|
||||
embeddings = embedding_dict.get(message_obj.public_id, [])
|
||||
for embedding in embeddings:
|
||||
for chunk_index, embedding in enumerate(embeddings):
|
||||
# Create MessageEmbedding record for metadata tracking
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
embedding=embedding,
|
||||
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)
|
||||
|
||||
# Add all embedding objects to the session
|
||||
# Create vector record for external vector store
|
||||
vector_id = f"{message_obj.public_id}_{chunk_index}"
|
||||
vector_records.append(
|
||||
VectorRecord(
|
||||
id=vector_id,
|
||||
embedding=embedding,
|
||||
metadata={
|
||||
"message_id": message_obj.public_id,
|
||||
"session_name": session_name,
|
||||
"peer_name": message_obj.peer_name,
|
||||
"chunk_index": chunk_index,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# Add all embedding metadata objects to the session
|
||||
if embedding_objects:
|
||||
db.add_all(embedding_objects)
|
||||
await db.commit()
|
||||
|
||||
# Upsert vectors to external vector store
|
||||
if vector_records:
|
||||
await vector_store.upsert_many(namespace, vector_records)
|
||||
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to generate message embeddings for %s messages in workspace %s and session %s.",
|
||||
|
|
|
|||
|
|
@ -23,9 +23,6 @@ from src.utils.representation import (
|
|||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Fetch extra documents to ensure we have enough after filtering
|
||||
FILTER_OVERSAMPLING_FACTOR = 3
|
||||
|
||||
|
||||
class RepresentationManager:
|
||||
"""Unified manager for representation and document queries."""
|
||||
|
|
@ -398,30 +395,31 @@ class RepresentationManager:
|
|||
observed=self.observed,
|
||||
query=query,
|
||||
max_distance=max_distance,
|
||||
top_k=count * FILTER_OVERSAMPLING_FACTOR,
|
||||
top_k=count,
|
||||
filters=self._build_filter_conditions(level),
|
||||
)
|
||||
|
||||
# Sort by creation time and return top count
|
||||
# Sort by creation time
|
||||
docs_sorted: list[models.Document] = sorted(
|
||||
list(documents), key=lambda x: x.created_at, reverse=True
|
||||
)
|
||||
return docs_sorted[:count]
|
||||
return docs_sorted
|
||||
|
||||
def _build_filter_conditions(
|
||||
self,
|
||||
level: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build complete filter conditions for document queries."""
|
||||
conditions: list[dict[str, Any]] = []
|
||||
"""
|
||||
Build filter conditions for document queries.
|
||||
|
||||
Returns a flat dict of key-value pairs for vector store filtering.
|
||||
"""
|
||||
filters: dict[str, Any] = {}
|
||||
|
||||
if level:
|
||||
conditions.append({"level": level})
|
||||
filters["level"] = level
|
||||
|
||||
if not conditions:
|
||||
return {}
|
||||
|
||||
return conditions[0] if len(conditions) == 1 else {"AND": conditions}
|
||||
return filters
|
||||
|
||||
|
||||
# Module-level functions for backward compatibility and convenience
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from src.cache.client import cache, get_cache_namespace
|
|||
from src.config import settings
|
||||
from src.exceptions import ConflictException, ResourceNotFoundException
|
||||
from src.utils.filter import apply_filter
|
||||
from src.vector_store import get_vector_store
|
||||
|
||||
logger = getLogger(__name__)
|
||||
|
||||
|
|
@ -251,6 +252,14 @@ async def delete_workspace(db: AsyncSession, workspace_name: str) -> schemas.Wor
|
|||
)
|
||||
)
|
||||
|
||||
# Get all collections for this workspace to delete their vector namespaces
|
||||
collections_result = await db.execute(
|
||||
select(models.Collection).where(
|
||||
models.Collection.workspace_name == workspace_name
|
||||
)
|
||||
)
|
||||
collections = collections_result.scalars().all()
|
||||
|
||||
await db.execute(
|
||||
delete(models.MessageEmbedding).where(
|
||||
models.MessageEmbedding.workspace_name == workspace_name
|
||||
|
|
@ -293,6 +302,45 @@ async def delete_workspace(db: AsyncSession, workspace_name: str) -> schemas.Wor
|
|||
await db.delete(honcho_workspace)
|
||||
await db.commit()
|
||||
|
||||
# Delete vector store namespaces for this workspace
|
||||
vector_store = get_vector_store()
|
||||
|
||||
# Delete message embeddings namespace for this workspace
|
||||
message_namespace = vector_store.get_message_namespace(workspace_name)
|
||||
try:
|
||||
await vector_store.delete_namespace(message_namespace)
|
||||
logger.debug(
|
||||
"Deleted message embeddings namespace %s for workspace %s",
|
||||
message_namespace,
|
||||
workspace_name,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to delete message embeddings namespace %s: %s",
|
||||
message_namespace,
|
||||
e,
|
||||
)
|
||||
|
||||
# Delete document embeddings namespaces for each collection
|
||||
for collection in collections:
|
||||
doc_namespace = vector_store.get_document_namespace(
|
||||
workspace_name, collection.observer, collection.observed
|
||||
)
|
||||
try:
|
||||
await vector_store.delete_namespace(doc_namespace)
|
||||
logger.debug(
|
||||
"Deleted document namespace %s for collection %s/%s",
|
||||
doc_namespace,
|
||||
collection.observer,
|
||||
collection.observed,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to delete document namespace %s: %s",
|
||||
doc_namespace,
|
||||
e,
|
||||
)
|
||||
|
||||
cache_key = workspace_cache_key(workspace_name)
|
||||
workspace_pattern = f"{cache_key}*"
|
||||
await cache.delete_match(workspace_pattern)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ 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,
|
||||
|
|
@ -22,7 +21,6 @@ from sqlalchemy import (
|
|||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB, TEXT
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.orm.properties import MappedColumn
|
||||
from sqlalchemy.sql import func
|
||||
from typing_extensions import override
|
||||
|
||||
|
|
@ -273,13 +271,21 @@ 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))
|
||||
message_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("messages.public_id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
|
|
@ -291,6 +297,8 @@ 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)
|
||||
|
||||
__table_args__ = (
|
||||
# Compound foreign key constraints
|
||||
|
|
@ -302,14 +310,6 @@ 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,6 +359,14 @@ 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(
|
||||
|
|
@ -371,7 +379,6 @@ class Document(Base):
|
|||
times_derived: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, server_default=text("1")
|
||||
)
|
||||
embedding: MappedColumn[Any] = mapped_column(Vector(1536))
|
||||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), index=True
|
||||
)
|
||||
|
|
@ -412,16 +419,6 @@ 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
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from src.config import settings
|
|||
from src.embedding_client import embedding_client
|
||||
from src.exceptions import ValidationException
|
||||
from src.utils.filter import apply_filter
|
||||
from src.vector_store import get_vector_store
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
|
@ -65,17 +66,19 @@ def reciprocal_rank_fusion(*ranked_lists: list[T], k: int = 60, limit: int) -> l
|
|||
async def _semantic_search(
|
||||
db: AsyncSession,
|
||||
query: str,
|
||||
stmt: Select[tuple[models.Message]],
|
||||
workspace_name: str,
|
||||
limit: int,
|
||||
filters: dict[str, Any] | None = None,
|
||||
) -> list[models.Message]:
|
||||
"""
|
||||
Perform semantic search using message embeddings.
|
||||
Perform semantic search using external vector store for message embeddings.
|
||||
|
||||
Args:
|
||||
db: Database session
|
||||
query: Search query
|
||||
stmt: Base SQL query conditions
|
||||
workspace_name: Name of the workspace to search in
|
||||
limit: Maximum number of results to return
|
||||
filters: Optional filters to apply at vector store level (supports: session_id, peer_id)
|
||||
|
||||
Returns:
|
||||
list of messages ordered by semantic similarity
|
||||
|
|
@ -87,16 +90,61 @@ async def _semantic_search(
|
|||
f"Query exceeds maximum token limit of {settings.MAX_EMBEDDING_TOKENS}."
|
||||
) from e
|
||||
|
||||
# Use cosine distance for semantic search on MessageEmbedding table
|
||||
semantic_query = stmt.join(
|
||||
models.MessageEmbedding,
|
||||
models.Message.public_id == models.MessageEmbedding.message_id,
|
||||
).order_by(models.MessageEmbedding.embedding.cosine_distance(embedding_query))
|
||||
# Get vector store and namespace for this workspace's messages
|
||||
vector_store = get_vector_store()
|
||||
namespace = vector_store.get_message_namespace(workspace_name)
|
||||
|
||||
semantic_query = semantic_query.limit(limit)
|
||||
# Build vector store filters from the provided filters
|
||||
vector_filters: dict[str, Any] = {}
|
||||
if filters:
|
||||
# Map external filter keys to vector store metadata keys
|
||||
if "session_id" in filters:
|
||||
vector_filters["session_name"] = filters["session_id"]
|
||||
if "peer_id" in filters:
|
||||
vector_filters["peer_name"] = filters["peer_id"]
|
||||
|
||||
# Query vector store for similar message embeddings
|
||||
# Since all filters are applied at the vector store level, we don't need to oversample
|
||||
vector_results = await vector_store.query(
|
||||
namespace,
|
||||
embedding_query,
|
||||
top_k=limit,
|
||||
filters=vector_filters if vector_filters else None,
|
||||
)
|
||||
|
||||
if not vector_results:
|
||||
return []
|
||||
|
||||
# Extract message IDs from vector results (vector ID format: {message_public_id}_{chunk_index})
|
||||
# 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_ids = list(seen_message_ids.keys())
|
||||
|
||||
# Fetch messages from database by the IDs from vector search
|
||||
# No additional filtering needed since vector store already applied all filters
|
||||
semantic_query = select(models.Message).where(
|
||||
models.Message.public_id.in_(message_ids)
|
||||
)
|
||||
|
||||
result = await db.execute(semantic_query)
|
||||
return list(result.scalars().all())
|
||||
messages = {msg.public_id: msg for msg in result.scalars().all()}
|
||||
|
||||
# Return messages in order of similarity (preserving vector store order)
|
||||
ordered_messages: list[models.Message] = []
|
||||
for msg_id in message_ids:
|
||||
if msg_id in messages:
|
||||
ordered_messages.append(messages[msg_id])
|
||||
|
||||
return ordered_messages
|
||||
|
||||
|
||||
async def _fulltext_search(
|
||||
|
|
@ -173,7 +221,7 @@ async def search(
|
|||
Args:
|
||||
db: Database session
|
||||
query: Search query to match against message content
|
||||
filters: Optional filters to scope search
|
||||
filters: Optional filters to scope search (must include workspace_id for semantic search)
|
||||
limit: Maximum number of results to return
|
||||
|
||||
Returns:
|
||||
|
|
@ -188,12 +236,18 @@ async def search(
|
|||
|
||||
search_results: list[list[models.Message]] = []
|
||||
|
||||
# Perform semantic search if enabled
|
||||
if settings.EMBED_MESSAGES:
|
||||
# Perform semantic search if enabled and we have workspace context
|
||||
# workspace_id is required for semantic search to determine the vector namespace
|
||||
workspace_name = filters.get("workspace_id") if filters else None
|
||||
if settings.EMBED_MESSAGES and workspace_name:
|
||||
# Get more results for fusion
|
||||
semantic_limit = limit * 2
|
||||
semantic_results = await _semantic_search(
|
||||
db=db, query=query, stmt=stmt, limit=semantic_limit
|
||||
db=db,
|
||||
query=query,
|
||||
workspace_name=workspace_name,
|
||||
limit=semantic_limit,
|
||||
filters=filters,
|
||||
)
|
||||
search_results.append(semantic_results)
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,194 @@
|
|||
"""
|
||||
Vector store abstraction layer for Honcho.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from src.config import settings
|
||||
from src.vector_store.lancedb import LanceDBVectorStore
|
||||
from src.vector_store.turbopuffer import TurbopufferVectorStore
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorRecord:
|
||||
"""A single vector record to be stored in the vector store."""
|
||||
|
||||
id: str
|
||||
embedding: list[float]
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QueryResult:
|
||||
"""A single result from a vector query."""
|
||||
|
||||
id: str
|
||||
score: float # Distance/similarity score (lower = more similar for cosine distance)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class VectorStore(ABC):
|
||||
"""
|
||||
Abstract base class for vector store implementations.
|
||||
|
||||
All vector operations are namespace-scoped. Namespaces map to:
|
||||
- Document embeddings: {prefix}-{workspace}-{observer}-{observed} (per collection)
|
||||
- Message embeddings: {prefix}-{workspace}-messages (per workspace)
|
||||
"""
|
||||
|
||||
namespace_prefix: str
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize the vector store.
|
||||
"""
|
||||
self.namespace_prefix = settings.VECTOR_STORE.NAMESPACE
|
||||
|
||||
# === Namespace helpers ===
|
||||
def get_document_namespace(
|
||||
self, workspace_name: str, observer: str, observed: str
|
||||
) -> str:
|
||||
"""
|
||||
Get the namespace for document embeddings (per collection).
|
||||
|
||||
Args:
|
||||
workspace_name: Name of the workspace
|
||||
observer: Name of the observing peer
|
||||
observed: Name of the observed peer
|
||||
|
||||
Returns:
|
||||
Namespace string in format: {prefix}_{workspace}_{observer}_{observed}
|
||||
"""
|
||||
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"
|
||||
|
||||
# === 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
|
||||
id: Unique identifier for the vector
|
||||
embedding: The embedding vector
|
||||
metadata: Optional metadata to store with the vector
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_many(
|
||||
self,
|
||||
namespace: str,
|
||||
vectors: list[VectorRecord],
|
||||
) -> None:
|
||||
"""
|
||||
Upsert multiple vectors into the store.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vectors in
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
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[QueryResult]:
|
||||
"""
|
||||
Query for similar vectors.
|
||||
|
||||
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 QueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def delete_many(self, namespace: str, ids: list[str]) -> None:
|
||||
"""
|
||||
Delete multiple vectors from the store.
|
||||
|
||||
Args:
|
||||
namespace: The namespace containing the vectors
|
||||
ids: List of vector identifiers to delete
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def delete_namespace(self, namespace: str) -> None:
|
||||
"""
|
||||
Delete an entire namespace and all its vectors.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to delete
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_vector_store_instance: VectorStore | None = None
|
||||
|
||||
|
||||
def get_vector_store() -> VectorStore:
|
||||
"""
|
||||
Get the configured vector store instance (singleton).
|
||||
|
||||
Returns:
|
||||
The vector store instance based on configuration.
|
||||
|
||||
Raises:
|
||||
ValueError: If the configured vector store type is invalid.
|
||||
"""
|
||||
global _vector_store_instance
|
||||
|
||||
if _vector_store_instance is not None:
|
||||
return _vector_store_instance
|
||||
|
||||
store_type = settings.VECTOR_STORE.TYPE
|
||||
|
||||
if store_type == "turbopuffer":
|
||||
_vector_store_instance = TurbopufferVectorStore()
|
||||
elif store_type == "lancedb":
|
||||
_vector_store_instance = LanceDBVectorStore()
|
||||
else:
|
||||
raise ValueError(f"Unknown vector store type: {store_type}")
|
||||
|
||||
return _vector_store_instance
|
||||
|
||||
|
||||
__all__ = [
|
||||
"VectorStore",
|
||||
"VectorRecord",
|
||||
"QueryResult",
|
||||
"get_vector_store",
|
||||
]
|
||||
|
|
@ -0,0 +1,287 @@
|
|||
"""
|
||||
LanceDB vector store implementation.
|
||||
|
||||
This module provides a LanceDB-based implementation of the VectorStore interface
|
||||
for use in self-hosted deployments of Honcho.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import lancedb
|
||||
import pyarrow as pa
|
||||
|
||||
from src.config import settings
|
||||
|
||||
from . import QueryResult, VectorRecord, VectorStore
|
||||
|
||||
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
|
||||
|
||||
|
||||
class LanceDBVectorStore(VectorStore):
|
||||
"""
|
||||
LanceDB implementation of the VectorStore interface.
|
||||
|
||||
Uses LanceDB's embedded mode for local vector storage.
|
||||
Each namespace corresponds to a LanceDB table.
|
||||
"""
|
||||
|
||||
_db: lancedb.DBConnection
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the LanceDB vector store."""
|
||||
super().__init__()
|
||||
self._db = lancedb.connect(settings.VECTOR_STORE.LANCEDB_PATH)
|
||||
|
||||
def _get_table(self, namespace: str) -> lancedb.table.Table | None:
|
||||
"""Get a table if it exists, otherwise return None."""
|
||||
if namespace in self._db.table_names():
|
||||
return self._db.open_table(namespace)
|
||||
return None
|
||||
|
||||
def _get_or_create_table(
|
||||
self, namespace: str, sample_data: list[dict[str, Any]] | None = None
|
||||
) -> lancedb.table.Table:
|
||||
"""
|
||||
Get existing table or create if not exists.
|
||||
|
||||
Args:
|
||||
namespace: Table name (namespace)
|
||||
sample_data: Optional sample data to infer schema from
|
||||
|
||||
Returns:
|
||||
LanceDB table
|
||||
"""
|
||||
if namespace in self._db.table_names():
|
||||
return self._db.open_table(namespace)
|
||||
|
||||
# Create table with sample data if provided
|
||||
if sample_data:
|
||||
return self._db.create_table(namespace, data=sample_data)
|
||||
|
||||
# Create empty table with base schema
|
||||
schema = pa.schema(
|
||||
[
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("vector", pa.list_(pa.float32(), VECTOR_DIMENSION)),
|
||||
]
|
||||
)
|
||||
return self._db.create_table(namespace, schema=schema)
|
||||
|
||||
def _row_to_dict(self, vector: VectorRecord) -> dict[str, Any]:
|
||||
"""Convert a VectorRecord to a dict for LanceDB."""
|
||||
row: dict[str, Any] = {
|
||||
"id": vector.id,
|
||||
"vector": vector.embedding,
|
||||
}
|
||||
# Add metadata fields
|
||||
if vector.metadata:
|
||||
row.update(vector.metadata)
|
||||
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 = self._get_or_create_table(namespace, sample_data=[row])
|
||||
|
||||
# Use merge_insert for upsert behavior
|
||||
table.merge_insert("id").when_matched_update_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:
|
||||
"""
|
||||
Upsert multiple vectors into LanceDB.
|
||||
|
||||
Args:
|
||||
namespace: The namespace (table) to store the vectors in
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
if not vectors:
|
||||
return
|
||||
|
||||
try:
|
||||
rows = [self._row_to_dict(v) for v in vectors]
|
||||
table = self._get_or_create_table(namespace, sample_data=rows)
|
||||
|
||||
# Use merge_insert for upsert behavior
|
||||
table.merge_insert("id").when_matched_update_all().execute(rows)
|
||||
|
||||
logger.debug(f"Upserted {len(vectors)} vectors to namespace {namespace}")
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert {len(vectors)} vectors to namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
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[QueryResult]:
|
||||
"""
|
||||
Query for similar vectors in LanceDB.
|
||||
|
||||
Args:
|
||||
namespace: The namespace (table) 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 QueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
table = self._get_table(namespace)
|
||||
if table is None:
|
||||
logger.debug(f"Table {namespace} does not exist, returning empty results")
|
||||
return []
|
||||
|
||||
try:
|
||||
# Build query (LanceDB types are incomplete, so type checker reports false positive)
|
||||
query = table.search(embedding).distance_type("cosine").limit(top_k) # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
# Apply filters if provided
|
||||
if filters:
|
||||
where_clause = self._build_where_clause(filters)
|
||||
if where_clause:
|
||||
query = query.where(where_clause)
|
||||
|
||||
# Execute query
|
||||
results = query.to_list()
|
||||
|
||||
# Convert to QueryResult objects
|
||||
query_results: list[QueryResult] = []
|
||||
for row in results:
|
||||
dist = float(row.get("_distance", 0.0))
|
||||
|
||||
# Filter by max_distance if specified
|
||||
if max_distance is not None and dist > max_distance:
|
||||
continue
|
||||
|
||||
# Extract metadata (everything except id, vector, _distance)
|
||||
metadata: dict[str, Any] = {
|
||||
k: v
|
||||
for k, v in row.items()
|
||||
if k not in ("id", "vector", "_distance")
|
||||
}
|
||||
|
||||
query_results.append(
|
||||
QueryResult(
|
||||
id=str(row["id"]),
|
||||
score=dist,
|
||||
metadata=metadata,
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"Query returned {len(query_results)} results from namespace {namespace}"
|
||||
)
|
||||
return query_results
|
||||
|
||||
except Exception:
|
||||
logger.exception(f"Failed to query namespace {namespace}")
|
||||
raise
|
||||
|
||||
def _build_where_clause(self, filters: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
Convert a filter dict to SQL WHERE clause syntax.
|
||||
|
||||
Args:
|
||||
filters: Dictionary of attribute -> value filters
|
||||
|
||||
Returns:
|
||||
SQL WHERE clause string or None if no filters
|
||||
"""
|
||||
if not filters:
|
||||
return None
|
||||
|
||||
conditions: list[str] = []
|
||||
for key, value in filters.items():
|
||||
# Handle string values with proper quoting
|
||||
if isinstance(value, str):
|
||||
# Escape single quotes in the value
|
||||
escaped_value = value.replace("'", "''")
|
||||
conditions.append(f"{key} = '{escaped_value}'")
|
||||
elif isinstance(value, bool):
|
||||
conditions.append(f"{key} = {str(value).lower()}")
|
||||
elif value is None:
|
||||
conditions.append(f"{key} IS NULL")
|
||||
else:
|
||||
conditions.append(f"{key} = {value}")
|
||||
|
||||
return " AND ".join(conditions) if conditions else None
|
||||
|
||||
async def delete_many(self, namespace: str, ids: list[str]) -> None:
|
||||
"""
|
||||
Delete multiple vectors from LanceDB.
|
||||
|
||||
Args:
|
||||
namespace: The namespace (table) containing the vectors
|
||||
ids: List of vector identifiers to delete
|
||||
"""
|
||||
if not ids:
|
||||
return
|
||||
|
||||
table = self._get_table(namespace)
|
||||
if table is None:
|
||||
logger.debug(f"Table {namespace} does not exist, nothing to delete")
|
||||
return
|
||||
|
||||
try:
|
||||
# Build IN clause with properly escaped IDs
|
||||
escaped_ids = [f"'{id.replace(chr(39), chr(39) + chr(39))}'" for id in ids]
|
||||
in_clause = ", ".join(escaped_ids)
|
||||
table.delete(f"id IN ({in_clause})")
|
||||
logger.debug(f"Deleted {len(ids)} vectors from namespace {namespace}")
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to delete {len(ids)} vectors from namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
async def delete_namespace(self, namespace: str) -> None:
|
||||
"""
|
||||
Delete an entire namespace (table) and all its vectors from LanceDB.
|
||||
|
||||
Args:
|
||||
namespace: The namespace (table) to delete
|
||||
"""
|
||||
try:
|
||||
if namespace in self._db.table_names():
|
||||
self._db.drop_table(namespace)
|
||||
logger.debug(f"Deleted namespace {namespace}")
|
||||
else:
|
||||
logger.debug(f"Namespace {namespace} does not exist, nothing to delete")
|
||||
except Exception:
|
||||
logger.exception(f"Failed to delete namespace {namespace}")
|
||||
raise
|
||||
|
|
@ -0,0 +1,285 @@
|
|||
"""
|
||||
Turbopuffer vector store implementation.
|
||||
|
||||
This module provides a Turbopuffer-based implementation of the VectorStore interface
|
||||
for use in managed deployments of Honcho.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Literal
|
||||
|
||||
from turbopuffer import Turbopuffer
|
||||
from turbopuffer.lib.namespace import Namespace
|
||||
from turbopuffer.types import Filter
|
||||
|
||||
from src.config import settings
|
||||
|
||||
from . import QueryResult, VectorRecord, VectorStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Type alias for Turbopuffer's equality filter format
|
||||
EqFilter = tuple[str, Literal["Eq"], Any]
|
||||
AndFilter = tuple[Literal["And"], Sequence[Filter]]
|
||||
|
||||
DISTANCE_METRIC = "cosine_distance"
|
||||
|
||||
|
||||
class TurbopufferVectorStore(VectorStore):
|
||||
"""
|
||||
Turbopuffer implementation of the VectorStore interface.
|
||||
|
||||
Uses Turbopuffer's Python SDK for vector operations.
|
||||
Each namespace corresponds to either:
|
||||
- A document collection: {prefix}-{workspace}-{observer}-{observed}
|
||||
- A workspace's message embeddings: {prefix}-{workspace}-messages
|
||||
"""
|
||||
|
||||
tpuf: Turbopuffer
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize the Turbopuffer vector store.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
# Configure Turbopuffer client
|
||||
api_key = settings.VECTOR_STORE.TURBOPUFFER_API_KEY
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"VECTOR_STORE_TURBOPUFFER_API_KEY must be set for Turbopuffer vector store"
|
||||
)
|
||||
|
||||
# Initialize the Turbopuffer client
|
||||
# Region can be configured via VECTOR_STORE_TURBOPUFFER_REGION or TURBOPUFFER_REGION env var
|
||||
region = settings.VECTOR_STORE.TURBOPUFFER_REGION or "gcp-us-east4"
|
||||
self.tpuf = Turbopuffer(api_key=api_key, region=region)
|
||||
|
||||
def _get_namespace(self, namespace: str) -> Namespace:
|
||||
"""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
|
||||
id: Unique identifier for the vector
|
||||
embedding: The embedding vector
|
||||
metadata: Optional metadata to store with the vector
|
||||
"""
|
||||
ns = self._get_namespace(namespace)
|
||||
attributes = vector.metadata or {}
|
||||
|
||||
try:
|
||||
# Build row data
|
||||
row: dict[str, Any] = {
|
||||
"id": vector.id,
|
||||
"vector": vector.embedding,
|
||||
**attributes,
|
||||
}
|
||||
|
||||
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:
|
||||
"""
|
||||
Upsert multiple vectors into Turbopuffer.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to store the vectors in
|
||||
vectors: List of VectorRecord objects to upsert
|
||||
"""
|
||||
if not vectors:
|
||||
return
|
||||
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
rows: list[dict[str, Any]] = [
|
||||
{
|
||||
"id": v.id,
|
||||
"vector": v.embedding,
|
||||
**v.metadata,
|
||||
}
|
||||
for v in vectors
|
||||
]
|
||||
|
||||
try:
|
||||
ns.write(
|
||||
upsert_rows=rows,
|
||||
distance_metric=DISTANCE_METRIC,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to upsert {len(vectors)} vectors to namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
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[QueryResult]:
|
||||
"""
|
||||
Query for similar vectors in Turbopuffer.
|
||||
|
||||
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 QueryResult objects, ordered by similarity (most similar first)
|
||||
"""
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
try:
|
||||
# Build filter conditions for Turbopuffer
|
||||
filter_condition = self._build_filters(filters) if filters else None
|
||||
|
||||
# Query using rank_by for vector similarity
|
||||
# rank_by must be a tuple: (attribute, "ANN", vector)
|
||||
rank_by: tuple[str, Literal["ANN"], Sequence[float]] = (
|
||||
"vector",
|
||||
"ANN",
|
||||
embedding,
|
||||
)
|
||||
|
||||
# Only pass filters if we have them (avoid passing None)
|
||||
query_kwargs: dict[str, Any] = {
|
||||
"rank_by": rank_by,
|
||||
"top_k": top_k,
|
||||
"distance_metric": DISTANCE_METRIC,
|
||||
"include_attributes": True,
|
||||
}
|
||||
if filter_condition is not None:
|
||||
query_kwargs["filters"] = filter_condition
|
||||
|
||||
response = ns.query(**query_kwargs)
|
||||
|
||||
query_results: list[QueryResult] = []
|
||||
for row in response.rows or []:
|
||||
# Distance is accessed via row["$dist"]
|
||||
dist: float = float(row["$dist"]) if "$dist" in row else 0.0
|
||||
# Filter by max_distance if specified
|
||||
if max_distance is not None and dist > max_distance:
|
||||
continue
|
||||
|
||||
# Extract attributes from model_extra (excludes id, vector, $dist)
|
||||
row_metadata: dict[str, Any] = {}
|
||||
if row.model_extra:
|
||||
# Filter out internal fields like $dist
|
||||
row_metadata = {
|
||||
k: v
|
||||
for k, v in row.model_extra.items()
|
||||
if not k.startswith("$")
|
||||
}
|
||||
|
||||
query_results.append(
|
||||
QueryResult(
|
||||
id=str(row.id),
|
||||
score=dist,
|
||||
metadata=row_metadata,
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"Query returned {len(query_results)} results from namespace {namespace}"
|
||||
)
|
||||
return query_results
|
||||
|
||||
except Exception:
|
||||
logger.exception(f"Failed to query namespace {namespace}")
|
||||
raise
|
||||
|
||||
def _build_filters(self, filters: dict[str, Any]) -> Filter | None:
|
||||
"""
|
||||
Convert a filter dict to Turbopuffer filter format.
|
||||
|
||||
Turbopuffer uses tuples like (attribute, "Eq", value) for filters,
|
||||
and ("And", [filters]) for combining multiple filters.
|
||||
|
||||
Args:
|
||||
filters: Dictionary of attribute -> value filters
|
||||
|
||||
Returns:
|
||||
Turbopuffer Filter or None if no filters
|
||||
"""
|
||||
if not filters:
|
||||
return None
|
||||
|
||||
filter_list: list[EqFilter] = []
|
||||
for key, value in filters.items():
|
||||
# Simple equality filter using "Eq" operator
|
||||
eq_filter: EqFilter = (key, "Eq", value)
|
||||
filter_list.append(eq_filter)
|
||||
|
||||
if not filter_list:
|
||||
return None
|
||||
|
||||
if len(filter_list) == 1:
|
||||
return filter_list[0]
|
||||
|
||||
# Combine multiple filters with AND
|
||||
and_filter: AndFilter = ("And", filter_list)
|
||||
return and_filter
|
||||
|
||||
async def delete_many(self, namespace: str, ids: list[str]) -> None:
|
||||
"""
|
||||
Delete multiple vectors from Turbopuffer.
|
||||
|
||||
Args:
|
||||
namespace: The namespace containing the vectors
|
||||
ids: List of vector identifiers to delete
|
||||
"""
|
||||
if not ids:
|
||||
return
|
||||
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
try:
|
||||
ns.write(deletes=ids)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
f"Failed to delete {len(ids)} vectors from namespace {namespace}"
|
||||
)
|
||||
raise
|
||||
|
||||
async def delete_namespace(self, namespace: str) -> None:
|
||||
"""
|
||||
Delete an entire namespace and all its vectors from Turbopuffer.
|
||||
|
||||
Args:
|
||||
namespace: The namespace to delete
|
||||
"""
|
||||
ns = self._get_namespace(namespace)
|
||||
|
||||
try:
|
||||
ns.delete_all()
|
||||
logger.debug(f"Deleted all vectors from namespace {namespace}")
|
||||
except Exception:
|
||||
logger.exception(f"Failed to delete namespace {namespace}")
|
||||
raise
|
||||
Loading…
Reference in New Issue