feat(tenancy): add Workspace multi-tenant foundation (#2353)

* Document multi-tenant workspace architecture

* Add OSS and commercial workspace boundaries

* docs: redesign multi-tenant workspace architecture

* feat(tenancy): implement workspace isolation

* docs(tenancy): record verification evidence

* docs(tenancy): revise single-instance SaaS topology

* docs(tenancy): refine architecture options

* docs: finalize cloud v2 multi-tenant decisions

* feat(tenancy): establish cloud isolation foundations

* feat(tenancy): harden shared cloud runtime boundaries

* docs(tenancy): record final isolation verification

* fix(tenancy): close isolation and permission gaps

* docs(tenancy): record final isolation verification

* feat(tenancy): connect cloud workspace control plane

* fix(build): install git for pinned SDK

* docs(cloud): update control plane verification

* chore: update multi-tenant SDK pin

* fix(cloud): skip legacy model sync during startup

* test(cloud): preserve minimal model manager fixtures

* fix(cloud): preserve authenticated account context

* fix(cloud): reuse authenticated account for user info

* feat(cloud): complete Workspace settings navigation

* test(web): cover Workspace dropdown menu

* feat(web): place workspace controls in sidebar

* refactor(web): streamline workspace controls

* style(web): format workspace layout test

* fix(cloud): surface runtime and workspace plan status

* fix(plugin): keep runtime identity stable across restarts

* fix(ui): widen and center workspace switcher

* fix(ui): hide roles from workspace switcher

* fix(ui): align workspace switcher with sidebar entries

* feat(workspace): add in-product collaboration and direct Cloud launch

* style: format collaboration changes

* fix(workspace): bind collaboration APIs to tenant UoW

* fix(cloud): preserve Core-owned collaboration state

* test(cloud): require Space identity for invite registration

* feat(cloud): complete secure invitation experience

* style(web): format invitation flows

* fix(cloud): recover box runtime without unscoped skill reload

* feat(oss): enforce invitation account and owner billing flows

* style: format OSS account service

* test(oss): cover invitation logout handoff

* fix(oss): resolve workspace owner in scoped session

* feat(cloud): harden multi-tenant runtime resources

* fix(cloud): bound runtime restart storms

* fix(cloud): eliminate periodic runtime CPU spikes

* fix(cloud): enforce instance capacity ceilings

* fix(cloud): scope public login capability discovery

* fix(cloud): bound tenant maintenance and monitoring work

* fix(runtime): bound tenant resource amplification

* fix(deps): pin green multi-tenant plugin SDK

* fix(cloud): handle unavailable skill capability

* fix(security): require authentication for image file endpoint (H-2)

- Changed /api/v1/files/image from AuthType.NONE to USER_TOKEN_OR_API_KEY
- Added Permission.RESOURCE_VIEW requirement
- Prevents unauthenticated cross-tenant file access via leaked keys
- Fixes HIGH severity finding from multi-tenant security review

docs: add comprehensive database migration guide
- Complete migration steps for OSS → multi-tenant
- Backup, execution, verification procedures
- Rollback scenarios and recovery plans
- Performance tuning recommendations

* test: add comprehensive cross-tenant isolation tests

Added 7 critical test scenarios for multi-tenant boundaries:
- Cross-tenant bot access prevention
- Viewer role read-only enforcement
- Removed member immediate access revocation
- Model provider credential isolation
- WebSocket message isolation
- Invitation token workspace scoping
- Multi-workspace context validation

These tests address P0-2 coverage gaps for:
- workspaces.py (membership & invitation flows)
- user.py (authentication & authorization)
- websocket_chat.py (real-time isolation)
- plugins.py (resource access control)

docs: finalize database migration guide

* fix(security): resolve M-1, M-2, M-3 security findings

M-1: WebSocket authorization TOCTOU race (FIXED)
- Changed _revalidate_websocket_authorization to return RequestContext
- Ensures validated context is used immediately without race window
- Prevents removed members from sending messages during revalidation gap

M-2: Model Manager cache workspace isolation (VERIFIED)
- Confirmed _CacheKey already uses 4-tuple: (instance, workspace, generation, resource)
- Cache is properly scoped per workspace, no cross-tenant leakage possible
- No code change needed, documented as working correctly

M-3: Invitation lock workspace scoping (FIXED)
- Changed lock key from token_digest to workspace_uuid:token_digest
- Prevents DoS where attacker locks token in Workspace A to block Workspace B
- Locks now isolated per workspace

All MEDIUM severity findings from security review now resolved.

* fix(cloud): unblock tenant CI and enforce knowledge quotas

* fix(tenancy): scope rerank model sync

---------

Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
RockChinQ
2026-07-30 21:43:35 +08:00
committed by GitHub
parent 463b120923
commit e1ac5e0fc8
468 changed files with 78320 additions and 13137 deletions
+19 -3
View File
@@ -2,7 +2,12 @@ from __future__ import annotations
import asyncio
from typing import Any
from chromadb import PersistentClient
from langbot.pkg.vector.vdb import VectorDatabase, SearchType
from langbot.pkg.vector.vdb import (
VectorDatabase,
SearchType,
remember_bounded_mapping,
runtime_cache_limit,
)
from langbot.pkg.core import app
import chromadb
import chromadb.errors
@@ -16,6 +21,13 @@ class ChromaVectorDatabase(VectorDatabase):
self.ap = ap
self.client = PersistentClient(path=base_path)
self._collections = {}
self._runtime_cache_limit = runtime_cache_limit(ap)
async def close(self) -> None:
# Chroma's PersistentClient has no public close API. Collection
# wrappers are safe to discard and otherwise retain every collection
# touched during the lifetime of this application object.
self._collections.clear()
@classmethod
def supported_search_types(cls) -> list[SearchType]:
@@ -23,8 +35,12 @@ class ChromaVectorDatabase(VectorDatabase):
async def get_or_create_collection(self, collection: str) -> chromadb.Collection:
if collection not in self._collections:
self._collections[collection] = await asyncio.to_thread(
self.client.get_or_create_collection, name=collection
runtime_collection = await asyncio.to_thread(self.client.get_or_create_collection, name=collection)
remember_bounded_mapping(
self._collections,
collection,
runtime_collection,
self._runtime_cache_limit,
)
self.ap.logger.info(f"Chroma collection '{collection}' accessed/created.")
return self._collections[collection]
+10 -2
View File
@@ -3,9 +3,10 @@ import asyncio
from typing import Any, Dict
from pymilvus import MilvusClient, DataType, CollectionSchema, FieldSchema
from pymilvus.milvus_client.index import IndexParams
from langbot.pkg.vector.vdb import VectorDatabase
from langbot.pkg.vector.vdb import VectorDatabase, remember_bounded_set, runtime_cache_limit
from langbot.pkg.vector.filter_utils import normalize_filter, strip_unsupported_fields
from langbot.pkg.core import app
from langbot.pkg.utils import bounded_executor
# Milvus schema only stores these metadata fields; filter on other fields is
# silently dropped with a warning.
@@ -71,8 +72,15 @@ class MilvusVectorDatabase(VectorDatabase):
self.db_name = db_name
self.client = None
self._collections: set[str] = set()
self._runtime_cache_limit = runtime_cache_limit(ap)
self._initialize_client()
async def close(self) -> None:
self._collections.clear()
if self.client is not None:
await bounded_executor.run_blocking_cleanup(self.client.close)
self.client = None
def _initialize_client(self):
"""Initialize Milvus client connection"""
try:
@@ -169,7 +177,7 @@ class MilvusVectorDatabase(VectorDatabase):
await self._ensure_index_if_missing(collection)
self.ap.logger.info(f"Milvus collection '{collection}' already exists")
self._collections.add(collection)
remember_bounded_set(self._collections, collection, self._runtime_cache_limit)
return collection
async def _ensure_index_if_missing(self, collection: str) -> None:
+322 -276
View File
@@ -1,22 +1,31 @@
from __future__ import annotations
from typing import Any, Dict
from sqlalchemy import create_engine, text, Column, String, Text
from sqlalchemy.orm import declarative_base
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
import contextlib
import dataclasses
from collections.abc import AsyncIterator
from typing import Any
import sqlalchemy
from pgvector.sqlalchemy import Vector
from langbot.pkg.vector.vdb import VectorDatabase
from langbot.pkg.vector.filter_utils import normalize_filter, strip_unsupported_fields
from sqlalchemy.dialects.postgresql import insert as postgresql_insert
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import declarative_base
from langbot.pkg.core import app
from langbot.pkg.vector.filter_utils import normalize_filter, strip_unsupported_fields
from langbot.pkg.vector.vdb import VectorDatabase
Base = declarative_base()
DEFAULT_ALLOWED_DIMENSIONS = (384, 512, 768, 1024, 1536)
# pgvector schema only stores these metadata fields.
_PG_SUPPORTED_FIELDS = {'text', 'file_id', 'chunk_uuid'}
# Callers use canonical metadata key 'uuid' but pgvector stores it as 'chunk_uuid'.
_PG_FIELD_ALIASES = {'uuid': 'chunk_uuid'}
# Map schema field names to SQLAlchemy columns (resolved lazily from PgVectorEntry).
_PG_COLUMN_MAP = {
'text': 'text',
'file_id': 'file_id',
@@ -24,21 +33,50 @@ _PG_COLUMN_MAP = {
}
@dataclasses.dataclass(frozen=True, slots=True)
class PgVectorScope:
"""Trusted relational tenant key for one knowledge-base operation."""
workspace_uuid: str
knowledge_base_uuid: str
embedding_dimension: int | None = None
def __post_init__(self) -> None:
for field_name in ('workspace_uuid', 'knowledge_base_uuid'):
value = getattr(self, field_name)
if not isinstance(value, str) or not value.strip():
raise ValueError(f'{field_name} must not be empty')
object.__setattr__(self, field_name, value.strip())
dimension = self.embedding_dimension
if dimension is not None and (isinstance(dimension, bool) or not isinstance(dimension, int) or dimension <= 0):
raise ValueError('embedding_dimension must be a positive integer')
class PgVectorEntry(Base):
"""SQLAlchemy model for pgvector entries"""
"""Tenant-scoped pgvector row created only by release/OSS migrations."""
__tablename__ = 'langbot_vectors'
id = Column(String, primary_key=True)
collection = Column(String, index=True, nullable=False)
embedding = Column(Vector(1536)) # Default dimension, will be created dynamically
text = Column(Text)
file_id = Column(String, index=True)
chunk_uuid = Column(String)
workspace_uuid = sqlalchemy.Column(sqlalchemy.String(36), primary_key=True)
knowledge_base_uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
vector_id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
embedding_dimension = sqlalchemy.Column(sqlalchemy.Integer, nullable=False)
embedding = sqlalchemy.Column(Vector(), nullable=False)
text = sqlalchemy.Column(sqlalchemy.Text)
file_id = sqlalchemy.Column(sqlalchemy.String(255), index=True)
chunk_uuid = sqlalchemy.Column(sqlalchemy.String(255))
__table_args__ = (
sqlalchemy.CheckConstraint(
'vector_dims(embedding) = embedding_dimension',
name='ck_langbot_vectors_embedding_dimension',
),
)
def _build_pg_conditions(filter_dict: dict[str, Any]) -> list:
"""Translate canonical filter dict into a list of SQLAlchemy conditions."""
"""Translate canonical filter dict into SQLAlchemy conditions."""
triples = normalize_filter(filter_dict)
triples = strip_unsupported_fields(triples, _PG_SUPPORTED_FIELDS, _PG_FIELD_ALIASES)
@@ -65,83 +103,139 @@ def _build_pg_conditions(filter_dict: dict[str, Any]) -> list:
class PgVectorDatabase(VectorDatabase):
"""PostgreSQL with pgvector extension database implementation"""
"""PostgreSQL vector adapter with explicit Workspace/RLS scope.
Cloud reuses the business database engine and never performs DDL. OSS can
still opt into a standalone pgvector database; that compatibility mode may
create a fresh schema, but it uses the same explicit tenant keys.
"""
def __init__(
self,
ap: app.Application,
connection_string: str = None,
connection_string: str | None = None,
host: str = 'localhost',
port: int = 5432,
database: str = 'langbot',
user: str = 'postgres',
password: str = 'postgres',
):
"""Initialize pgvector database
Args:
ap: Application instance
connection_string: Full PostgreSQL connection string (overrides other params)
host: PostgreSQL host
port: PostgreSQL port
database: Database name
user: Database user
password: Database password
"""
*,
use_business_database: bool = False,
allowed_dimensions: list[int] | tuple[int, ...] = DEFAULT_ALLOWED_DIMENSIONS,
) -> None:
self.ap = ap
self.use_business_database = use_business_database
self.allowed_dimensions = self._normalize_allowed_dimensions(allowed_dimensions)
self.engine = None
self.async_engine = None
self.AsyncSessionLocal: async_sessionmaker[AsyncSession] | None = None
if use_business_database:
persistence_mgr = getattr(ap, 'persistence_mgr', None)
if persistence_mgr is None:
raise RuntimeError('Shared pgvector requires the initialized business persistence manager')
business_engine = persistence_mgr.get_db_engine()
if business_engine.dialect.name != 'postgresql':
raise RuntimeError('Shared pgvector requires the PostgreSQL business database')
self.async_engine = business_engine
self.ap.logger.info('Connected pgvector adapter to the shared PostgreSQL business database')
return
# Build connection string if not provided
if connection_string:
self.connection_string = connection_string
else:
self.connection_string = f'postgresql+psycopg://{user}:{password}@{host}:{port}/{database}'
self.async_connection_string = self.connection_string.replace('postgresql://', 'postgresql+asyncpg://').replace(
'postgresql+psycopg://', 'postgresql+asyncpg://'
)
self._initialize_standalone_db()
self.engine = None
self.async_engine = None
self.SessionLocal = None
self.AsyncSessionLocal = None
self._collections = set()
self._initialize_db()
@staticmethod
def _normalize_allowed_dimensions(dimensions: list[int] | tuple[int, ...]) -> frozenset[int]:
if not isinstance(dimensions, (list, tuple)) or not dimensions:
raise ValueError('pgvector allowed_dimensions must be a non-empty list')
if any(isinstance(item, bool) or not isinstance(item, int) or item <= 0 for item in dimensions):
raise ValueError('pgvector allowed_dimensions must contain positive integers')
unsupported = set(dimensions) - set(DEFAULT_ALLOWED_DIMENSIONS)
if unsupported:
raise ValueError(f'pgvector dimensions do not have release-created ANN indexes: {sorted(unsupported)}')
return frozenset(dimensions)
def _initialize_db(self):
"""Initialize database connection and create tables"""
try:
# Create async engine for async operations
self.async_engine = create_async_engine(self.async_connection_string, echo=False, pool_pre_ping=True)
self.AsyncSessionLocal = async_sessionmaker(self.async_engine, class_=AsyncSession, expire_on_commit=False)
def _initialize_standalone_db(self) -> None:
"""Initialize the explicit OSS external database compatibility path."""
# Create sync engine for table creation
sync_connection_string = self.connection_string.replace('postgresql+asyncpg://', 'postgresql+psycopg://')
self.engine = create_engine(sync_connection_string, echo=False)
from sqlalchemy import create_engine
# Create pgvector extension and tables
with self.engine.connect() as conn:
# Enable pgvector extension
conn.execute(text('CREATE EXTENSION IF NOT EXISTS vector'))
conn.commit()
self.async_engine = create_async_engine(self.async_connection_string, echo=False, pool_pre_ping=True)
self.AsyncSessionLocal = async_sessionmaker(self.async_engine, class_=AsyncSession, expire_on_commit=False)
sync_connection_string = self.connection_string.replace('postgresql+asyncpg://', 'postgresql+psycopg://')
self.engine = create_engine(sync_connection_string, echo=False)
# Create tables
Base.metadata.create_all(self.engine)
with self.engine.begin() as conn:
conn.execute(sqlalchemy.text('CREATE EXTENSION IF NOT EXISTS vector'))
existing_tables = set(sqlalchemy.inspect(conn).get_table_names())
if PgVectorEntry.__tablename__ in existing_tables:
columns = {
column['name'] for column in sqlalchemy.inspect(conn).get_columns(PgVectorEntry.__tablename__)
}
required = {
'workspace_uuid',
'knowledge_base_uuid',
'vector_id',
'embedding_dimension',
'embedding',
}
if not required.issubset(columns):
raise RuntimeError(
'The external pgvector database uses the legacy unscoped schema; '
'migrate it before enabling multi-tenant vector access'
)
Base.metadata.create_all(conn)
self.ap.logger.info('Connected to PostgreSQL with pgvector')
except Exception as e:
self.ap.logger.error(f'Failed to connect to PostgreSQL: {e}')
raise
self.ap.logger.info('Connected to standalone PostgreSQL pgvector database')
def _require_scope(self, scope: PgVectorScope | None, *, require_dimension: bool) -> PgVectorScope:
if not isinstance(scope, PgVectorScope):
raise ValueError('pgvector operations require a trusted PgVectorScope')
dimension = scope.embedding_dimension
if require_dimension and dimension is None:
raise ValueError('pgvector operation requires an embedding dimension')
if dimension is not None and dimension not in self.allowed_dimensions:
raise ValueError(f'Embedding dimension {dimension} is not enabled for this pgvector deployment')
return scope
@staticmethod
def _scope_conditions(scope: PgVectorScope) -> tuple[Any, Any]:
return (
PgVectorEntry.workspace_uuid == scope.workspace_uuid,
PgVectorEntry.knowledge_base_uuid == scope.knowledge_base_uuid,
)
@contextlib.asynccontextmanager
async def _session(self, scope: PgVectorScope) -> AsyncIterator[AsyncSession]:
admission = getattr(self.ap, 'deployment_admission', None)
if admission is not None:
admission.require_active()
if self.use_business_database:
async with self.ap.persistence_mgr.tenant_uow(scope.workspace_uuid) as uow:
yield uow.session
if admission is not None:
admission.require_active()
return
if self.AsyncSessionLocal is None: # pragma: no cover - constructor invariant
raise RuntimeError('Standalone pgvector session factory is unavailable')
async with self.AsyncSessionLocal() as session, session.begin():
yield session
if admission is not None:
admission.require_active()
async def get_or_create_collection(self, collection: str):
"""Get or create a collection (logical grouping in pgvector)
"""Retain the common adapter API; relational rows need no collection DDL."""
Args:
collection: Collection name (knowledge base UUID)
"""
# In pgvector, collections are logical - we just track them
if collection not in self._collections:
self._collections.add(collection)
self.ap.logger.info(f"Registered pgvector collection '{collection}'")
if not isinstance(collection, str) or not collection.strip():
raise ValueError('collection must not be empty')
return collection
async def add_embeddings(
@@ -151,38 +245,59 @@ class PgVectorDatabase(VectorDatabase):
embeddings_list: list[list[float]],
metadatas: list[dict[str, Any]],
documents: list[str] | None = None,
*,
scope: PgVectorScope | None = None,
) -> None:
"""Add vector embeddings to pgvector
Args:
collection: Collection name
ids: List of unique IDs for each vector
embeddings_list: List of embedding vectors
metadatas: List of metadata dictionaries
"""
scope = self._require_scope(scope, require_dimension=True)
await self.get_or_create_collection(collection)
if not ids:
return
if len(ids) != len(embeddings_list) or len(metadatas) != len(ids):
raise ValueError('pgvector ids, embeddings and metadata lengths must match')
if documents is not None and len(documents) != len(ids):
raise ValueError('pgvector documents length must match ids')
if len(set(ids)) != len(ids) or any(not isinstance(item, str) or not item.strip() for item in ids):
raise ValueError('pgvector vector IDs must be unique non-empty strings per upsert')
expected_dimension = scope.embedding_dimension
if any(len(embedding) != expected_dimension for embedding in embeddings_list):
raise ValueError(f'All embeddings must have the selected dimension {expected_dimension}')
async with self.AsyncSessionLocal() as session:
try:
for i, vector_id in enumerate(ids):
metadata = metadatas[i] if i < len(metadatas) else {}
values = []
for index, vector_id in enumerate(ids):
metadata = metadatas[index]
document = documents[index] if documents is not None else None
values.append(
{
'workspace_uuid': scope.workspace_uuid,
'knowledge_base_uuid': scope.knowledge_base_uuid,
'vector_id': vector_id.strip(),
'embedding_dimension': expected_dimension,
'embedding': embeddings_list[index],
'text': metadata.get('text', document or ''),
'file_id': metadata.get('file_id', ''),
'chunk_uuid': metadata.get('uuid', metadata.get('chunk_uuid', '')),
}
)
entry = PgVectorEntry(
id=vector_id,
collection=collection,
embedding=embeddings_list[i],
text=metadata.get('text', ''),
file_id=metadata.get('file_id', ''),
chunk_uuid=metadata.get('uuid', ''),
)
session.add(entry)
await session.commit()
self.ap.logger.info(f"Added {len(ids)} embeddings to pgvector collection '{collection}'")
except Exception as e:
await session.rollback()
self.ap.logger.error(f'Error adding embeddings to pgvector: {e}')
raise
statement = postgresql_insert(PgVectorEntry).values(values)
excluded = statement.excluded
statement = statement.on_conflict_do_update(
index_elements=[
PgVectorEntry.workspace_uuid,
PgVectorEntry.knowledge_base_uuid,
PgVectorEntry.vector_id,
],
set_={
'embedding_dimension': excluded.embedding_dimension,
'embedding': excluded.embedding,
'text': excluded.text,
'file_id': excluded.file_id,
'chunk_uuid': excluded.chunk_uuid,
},
)
async with self._session(scope) as session:
await session.execute(statement)
self.ap.logger.info(f'Upserted {len(ids)} pgvector embeddings for knowledge base {scope.knowledge_base_uuid}')
async def search(
self,
@@ -193,125 +308,79 @@ class PgVectorDatabase(VectorDatabase):
query_text: str = '',
filter: dict[str, Any] | None = None,
vector_weight: float | None = None,
) -> Dict[str, Any]:
"""Search for similar vectors using cosine distance
Args:
collection: Collection name
query_embedding: Query vector
k: Number of top results to return
Returns:
Dictionary with search results in Chroma-compatible format
"""
*,
scope: PgVectorScope | None = None,
) -> dict[str, Any]:
del query_text, vector_weight
scope = self._require_scope(scope, require_dimension=True)
await self.get_or_create_collection(collection)
if search_type != 'vector':
raise ValueError('pgvector currently supports vector search only')
if k <= 0:
raise ValueError('pgvector search limit must be positive')
if len(query_embedding) != scope.embedding_dimension:
raise ValueError(f'Query embedding must have the selected dimension {scope.embedding_dimension}')
async with self.AsyncSessionLocal() as session:
try:
# Use cosine distance for similarity search
from sqlalchemy import select
typed_embedding = sqlalchemy.cast(PgVectorEntry.embedding, Vector(scope.embedding_dimension))
distance = typed_embedding.cosine_distance(query_embedding)
statement = (
sqlalchemy.select(
PgVectorEntry.vector_id,
PgVectorEntry.text,
PgVectorEntry.file_id,
PgVectorEntry.chunk_uuid,
distance.label('distance'),
)
.where(*self._scope_conditions(scope), PgVectorEntry.embedding_dimension == scope.embedding_dimension)
.order_by(distance)
.limit(k)
)
for condition in _build_pg_conditions(filter or {}):
statement = statement.where(condition)
# Query for similar vectors
stmt = (
select(
PgVectorEntry.id,
PgVectorEntry.text,
PgVectorEntry.file_id,
PgVectorEntry.chunk_uuid,
PgVectorEntry.embedding.cosine_distance(query_embedding).label('distance'),
)
.filter(PgVectorEntry.collection == collection)
.order_by(PgVectorEntry.embedding.cosine_distance(query_embedding))
.limit(k)
)
async with self._session(scope) as session:
rows = (await session.execute(statement)).all()
if filter:
for cond in _build_pg_conditions(filter):
stmt = stmt.filter(cond)
ids = [row.vector_id for row in rows]
distances = [float(row.distance) for row in rows]
metadatas = [
{'text': row.text or '', 'file_id': row.file_id or '', 'uuid': row.chunk_uuid or ''} for row in rows
]
return {'ids': [ids], 'distances': [distances], 'metadatas': [metadatas]}
result = await session.execute(stmt)
rows = result.fetchall()
# Convert to Chroma-compatible format
ids = []
distances = []
metadatas = []
for row in rows:
ids.append(row.id)
distances.append(float(row.distance))
metadatas.append(
{'text': row.text or '', 'file_id': row.file_id or '', 'uuid': row.chunk_uuid or ''}
)
result_dict = {'ids': [ids], 'distances': [distances], 'metadatas': [metadatas]}
self.ap.logger.info(f"pgvector search in '{collection}' returned {len(ids)} results")
return result_dict
except Exception as e:
self.ap.logger.error(f'Error searching pgvector: {e}')
raise
async def delete_by_file_id(self, collection: str, file_id: str) -> None:
"""Delete vectors by file_id
Args:
collection: Collection name
file_id: File ID to filter deletion
"""
async def delete_by_file_id(
self,
collection: str,
file_id: str,
*,
scope: PgVectorScope | None = None,
) -> None:
scope = self._require_scope(scope, require_dimension=False)
await self.get_or_create_collection(collection)
statement = sqlalchemy.delete(PgVectorEntry).where(
*self._scope_conditions(scope),
PgVectorEntry.file_id == file_id,
)
async with self._session(scope) as session:
await session.execute(statement)
async with self.AsyncSessionLocal() as session:
try:
from sqlalchemy import delete
stmt = delete(PgVectorEntry).where(
PgVectorEntry.collection == collection, PgVectorEntry.file_id == file_id
)
await session.execute(stmt)
await session.commit()
self.ap.logger.info(
f"Deleted embeddings from pgvector collection '{collection}' with file_id: {file_id}"
)
except Exception as e:
await session.rollback()
self.ap.logger.error(f'Error deleting from pgvector: {e}')
raise
async def delete_by_filter(self, collection: str, filter: dict[str, Any]) -> int:
"""Delete vectors matching a metadata filter.
Args:
collection: Collection name
filter: Canonical metadata filter dict
"""
async def delete_by_filter(
self,
collection: str,
filter: dict[str, Any],
*,
scope: PgVectorScope | None = None,
) -> int:
scope = self._require_scope(scope, require_dimension=False)
await self.get_or_create_collection(collection)
conditions = _build_pg_conditions(filter)
if not conditions:
self.ap.logger.warning(
f"pgvector delete_by_filter on '{collection}': filter produced no conditions, skipping"
)
self.ap.logger.warning('pgvector delete_by_filter produced no supported conditions; skipping')
return 0
await self.get_or_create_collection(collection)
async with self.AsyncSessionLocal() as session:
try:
from sqlalchemy import delete
stmt = delete(PgVectorEntry).where(PgVectorEntry.collection == collection)
for cond in conditions:
stmt = stmt.where(cond)
result = await session.execute(stmt)
await session.commit()
deleted = result.rowcount
self.ap.logger.info(f"Deleted {deleted} embeddings from pgvector collection '{collection}' by filter")
return deleted
except Exception as e:
await session.rollback()
self.ap.logger.error(f'Error deleting from pgvector by filter: {e}')
raise
statement = sqlalchemy.delete(PgVectorEntry).where(*self._scope_conditions(scope), *conditions)
async with self._session(scope) as session:
result = await session.execute(statement)
return int(result.rowcount or 0)
async def list_by_filter(
self,
@@ -319,85 +388,62 @@ class PgVectorDatabase(VectorDatabase):
filter: dict[str, Any] | None = None,
limit: int = 20,
offset: int = 0,
*,
scope: PgVectorScope | None = None,
) -> tuple[list[dict[str, Any]], int]:
scope = self._require_scope(scope, require_dimension=False)
await self.get_or_create_collection(collection)
if limit <= 0 or offset < 0:
raise ValueError('pgvector pagination requires limit > 0 and offset >= 0')
async with self.AsyncSessionLocal() as session:
try:
from sqlalchemy import select, func
conditions = [*self._scope_conditions(scope), *_build_pg_conditions(filter or {})]
statement = (
sqlalchemy.select(
PgVectorEntry.vector_id,
PgVectorEntry.text,
PgVectorEntry.file_id,
PgVectorEntry.chunk_uuid,
)
.where(*conditions)
.order_by(PgVectorEntry.vector_id)
.offset(offset)
.limit(limit)
)
count_statement = sqlalchemy.select(sqlalchemy.func.count()).select_from(PgVectorEntry).where(*conditions)
async with self._session(scope) as session:
rows = (await session.execute(statement)).all()
total = int((await session.execute(count_statement)).scalar_one())
stmt = (
select(
PgVectorEntry.id,
PgVectorEntry.text,
PgVectorEntry.file_id,
PgVectorEntry.chunk_uuid,
)
.filter(PgVectorEntry.collection == collection)
.offset(offset)
.limit(limit)
)
return (
[
{
'id': row.vector_id,
'document': row.text or '',
'metadata': {
'text': row.text or '',
'file_id': row.file_id or '',
'uuid': row.chunk_uuid or '',
},
}
for row in rows
],
total,
)
count_stmt = (
select(func.count()).select_from(PgVectorEntry).filter(PgVectorEntry.collection == collection)
)
async def delete_collection(
self,
collection: str,
*,
scope: PgVectorScope | None = None,
) -> None:
scope = self._require_scope(scope, require_dimension=False)
await self.get_or_create_collection(collection)
statement = sqlalchemy.delete(PgVectorEntry).where(*self._scope_conditions(scope))
async with self._session(scope) as session:
await session.execute(statement)
if filter:
for cond in _build_pg_conditions(filter):
stmt = stmt.filter(cond)
count_stmt = count_stmt.filter(cond)
result = await session.execute(stmt)
rows = result.fetchall()
count_result = await session.execute(count_stmt)
total = count_result.scalar() or 0
items = []
for row in rows:
items.append(
{
'id': row.id,
'document': row.text or '',
'metadata': {
'text': row.text or '',
'file_id': row.file_id or '',
'uuid': row.chunk_uuid or '',
},
}
)
return items, total
except Exception as e:
self.ap.logger.error(f'Error listing from pgvector: {e}')
raise
async def delete_collection(self, collection: str):
"""Delete all vectors in a collection
Args:
collection: Collection name to delete
"""
if collection in self._collections:
self._collections.remove(collection)
async with self.AsyncSessionLocal() as session:
try:
from sqlalchemy import delete
stmt = delete(PgVectorEntry).where(PgVectorEntry.collection == collection)
await session.execute(stmt)
await session.commit()
self.ap.logger.info(f"Deleted pgvector collection '{collection}'")
except Exception as e:
await session.rollback()
self.ap.logger.error(f'Error deleting pgvector collection: {e}')
raise
async def close(self):
"""Close database connections"""
if self.async_engine:
async def close(self) -> None:
if not self.use_business_database and self.async_engine is not None:
await self.async_engine.dispose()
if self.engine:
if self.engine is not None:
self.engine.dispose()
+8 -3
View File
@@ -4,7 +4,7 @@ from typing import Any, Dict, List
from qdrant_client import AsyncQdrantClient, models
from langbot.pkg.core import app
from langbot.pkg.vector.vdb import VectorDatabase
from langbot.pkg.vector.vdb import VectorDatabase, remember_bounded_set, runtime_cache_limit
from langbot.pkg.vector.filter_utils import normalize_filter
@@ -52,6 +52,11 @@ class QdrantVectorDatabase(VectorDatabase):
self.client = AsyncQdrantClient(host=host, port=int(port), api_key=api_key)
self._collections: set[str] = set()
self._runtime_cache_limit = runtime_cache_limit(ap)
async def close(self) -> None:
self._collections.clear()
await self.client.close()
async def _ensure_collection(self, collection: str, vector_size: int) -> None:
if collection in self._collections:
@@ -59,14 +64,14 @@ class QdrantVectorDatabase(VectorDatabase):
exists = await self.client.collection_exists(collection)
if exists:
self._collections.add(collection)
remember_bounded_set(self._collections, collection, self._runtime_cache_limit)
return
await self.client.create_collection(
collection_name=collection,
vectors_config=models.VectorParams(size=vector_size, distance=models.Distance.COSINE),
)
self._collections.add(collection)
remember_bounded_set(self._collections, collection, self._runtime_cache_limit)
self.ap.logger.info(f"Qdrant collection '{collection}' created with dim={vector_size}.")
async def get_or_create_collection(self, collection: str):
+57 -8
View File
@@ -7,7 +7,13 @@ from typing import Any, Dict, List
from langbot.pkg.core import app
from langbot.pkg.vector.vdb import VectorDatabase, SearchType
from langbot.pkg.utils import bounded_executor
from langbot.pkg.vector.vdb import (
VectorDatabase,
SearchType,
remember_bounded_mapping,
runtime_cache_limit,
)
try:
import pyseekdb
@@ -90,6 +96,7 @@ class SeekDBVectorDatabase(VectorDatabase):
self._collections: Dict[str, Any] = {}
self._collection_configs: Dict[str, HNSWConfiguration] = {}
self._runtime_cache_limit = runtime_cache_limit(ap)
self._escape_table = str.maketrans(
{
@@ -103,6 +110,13 @@ class SeekDBVectorDatabase(VectorDatabase):
}
)
async def close(self) -> None:
self._collections.clear()
self._collection_configs.clear()
close = getattr(self.client, 'close', None)
if callable(close):
await bounded_executor.run_blocking_cleanup(close)
def _normalize_collection_name(self, collection: str) -> str:
"""SeekDB only accepts [a-zA-Z0-9_], while LangBot uses UUID-like KB IDs."""
normalized = re.sub(r'[^A-Za-z0-9_]', '_', collection)
@@ -132,7 +146,12 @@ class SeekDBVectorDatabase(VectorDatabase):
if await asyncio.to_thread(self.client.has_collection, collection):
# Collection exists, get it
coll = await asyncio.to_thread(self.client.get_collection, collection, embedding_function=None)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
self.ap.logger.info(f"SeekDB collection '{collection}' retrieved.")
return coll
@@ -145,7 +164,12 @@ class SeekDBVectorDatabase(VectorDatabase):
# Create HNSW configuration
config = HNSWConfiguration(dimension=vector_size, distance='cosine')
self._collection_configs[collection] = config
remember_bounded_mapping(
self._collection_configs,
collection,
config,
self._runtime_cache_limit,
)
# Create collection without embedding function (we manage embeddings externally)
coll = await asyncio.to_thread(
@@ -155,7 +179,12 @@ class SeekDBVectorDatabase(VectorDatabase):
embedding_function=None, # Disable automatic embedding
)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
self.ap.logger.info(f"SeekDB collection '{collection}' created with dimension={vector_size}, distance='cosine'")
return coll
@@ -243,7 +272,12 @@ class SeekDBVectorDatabase(VectorDatabase):
# Get collection
if collection not in self._collections:
coll = await asyncio.to_thread(self.client.get_collection, collection, embedding_function=None)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
else:
coll = self._collections[collection]
@@ -349,7 +383,12 @@ class SeekDBVectorDatabase(VectorDatabase):
# Get collection
if collection not in self._collections:
coll = await asyncio.to_thread(self.client.get_collection, collection, embedding_function=None)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
else:
coll = self._collections[collection]
@@ -374,7 +413,12 @@ class SeekDBVectorDatabase(VectorDatabase):
if collection not in self._collections:
coll = await asyncio.to_thread(self.client.get_collection, collection, embedding_function=None)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
else:
coll = self._collections[collection]
@@ -396,7 +440,12 @@ class SeekDBVectorDatabase(VectorDatabase):
if collection not in self._collections:
coll = await asyncio.to_thread(self.client.get_collection, collection, embedding_function=None)
self._collections[collection] = coll
remember_bounded_mapping(
self._collections,
collection,
coll,
self._runtime_cache_limit,
)
else:
coll = self._collections[collection]
+39 -38
View File
@@ -6,7 +6,12 @@ import struct
from typing import Any
from langbot.pkg.core import app
from langbot.pkg.vector.vdb import VectorDatabase, SearchType
from langbot.pkg.vector.vdb import (
VectorDatabase,
SearchType,
remember_bounded_set,
runtime_cache_limit,
)
from langbot.pkg.vector.filter_utils import normalize_filter, strip_unsupported_fields
try:
@@ -77,6 +82,7 @@ _MATCH_ALL = '-@file_id:{__langbot_match_all_sentinel__}'
# files/filters matching more than one page of chunks are fully removed
# (no silent truncation / orphaned vectors).
_DELETE_SCAN_BATCH = 10000
_MAX_DELETE_SCAN_ROUNDS = 1000
# Characters Valkey Search's TAG query parser cannot handle even when
# backslash-escaped (the brace delimiters and the wildcard). file_id TAG
@@ -153,6 +159,7 @@ class ValkeySearchVectorDatabase(VectorDatabase):
self._client_lock = asyncio.Lock()
# Index names we have already ensured this process lifetime.
self._ensured_indexes: set[str] = set()
self._runtime_cache_limit = runtime_cache_limit(ap)
# Whether we have already warned about the non-honored vector_weight.
self._vector_weight_warned = False
@@ -364,7 +371,7 @@ class ValkeySearchVectorDatabase(VectorDatabase):
# check-then-create TOCTOU window.
try:
await ft.info(client, index)
self._ensured_indexes.add(index)
remember_bounded_set(self._ensured_indexes, index, self._runtime_cache_limit)
return
except RequestError:
pass
@@ -389,7 +396,7 @@ class ValkeySearchVectorDatabase(VectorDatabase):
]
options = FtCreateOptions(data_type=DataType.HASH, prefixes=[self._key_prefix(collection)])
await ft.create(client, index, schema, options)
self._ensured_indexes.add(index)
remember_bounded_set(self._ensured_indexes, index, self._runtime_cache_limit)
self.ap.logger.info(
f"Valkey Search index '{index}' created (dim={dim}, algo={self._algorithm.value}, "
f'metric={self._distance_metric.value})'
@@ -646,11 +653,9 @@ class ValkeySearchVectorDatabase(VectorDatabase):
return
query = f'@{_FIELD_FILE_ID}:{{{self._encode_and_escape_tag(file_id)}}}'
keys = await self._search_keys(client, index, query)
if keys:
await client.delete(keys)
deleted = await self._delete_search_results(client, index, query)
self.ap.logger.info(
f"Deleted {len(keys)} embeddings from Valkey Search collection '{collection}' with file_id: {file_id}"
f"Deleted {deleted} embeddings from Valkey Search collection '{collection}' with file_id: {file_id}"
)
async def delete_by_filter(self, collection: str, filter: dict[str, Any]) -> int:
@@ -670,11 +675,9 @@ class ValkeySearchVectorDatabase(VectorDatabase):
collection,
)
return 0
keys = await self._search_keys(client, index, query)
if keys:
await client.delete(keys)
self.ap.logger.info(f"Deleted {len(keys)} embeddings from Valkey Search collection '{collection}' by filter")
return len(keys)
deleted = await self._delete_search_results(client, index, query)
self.ap.logger.info(f"Deleted {deleted} embeddings from Valkey Search collection '{collection}' by filter")
return deleted
async def list_by_filter(
self,
@@ -772,43 +775,35 @@ class ValkeySearchVectorDatabase(VectorDatabase):
# was being paid on the first query to each collection.
try:
await ft.info(client, index)
self._ensured_indexes.add(index)
remember_bounded_set(self._ensured_indexes, index, self._runtime_cache_limit)
return True
except RequestError:
return False
async def _search_keys(self, client: GlideClient, index: str, query: str) -> list[str]:
"""Return all matching document keys for a query (NOCONTENT).
async def _delete_search_results(self, client: GlideClient, index: str, query: str) -> int:
"""Delete matching hashes in fixed batches without retaining every key.
Paginates through the full result set in pages of ``_DELETE_SCAN_BATCH``
so that queries matching more than one page of chunks are fully
enumerated (avoids silently truncating deletes and leaving orphaned
vectors).
Each deletion shrinks the result set, so every search starts at offset
zero. Advancing an offset after deleting the preceding page would skip
records as the remaining results shift left.
"""
keys: list[str] = []
offset = 0
while True:
deleted = 0
for _round in range(_MAX_DELETE_SCAN_ROUNDS):
options = FtSearchOptions(
nocontent=True,
limit=FtSearchLimit(offset, _DELETE_SCAN_BATCH),
limit=FtSearchLimit(0, _DELETE_SCAN_BATCH),
dialect=2,
)
try:
reply = await ft.search(client, index, query, options)
except Exception as exc:
if self._is_missing_index_error(exc):
return keys
return deleted
raise
if not reply or len(reply) < 2:
break
# reply[0] is the total match count; reply[1] holds this page.
total = 0
try:
total = int(reply[0])
except (TypeError, ValueError):
total = 0
return deleted
docs = reply[1]
if isinstance(docs, dict):
@@ -819,11 +814,17 @@ class ValkeySearchVectorDatabase(VectorDatabase):
page = []
if not page:
break
keys.extend(page)
return deleted
await client.delete(page)
deleted += len(page)
offset += len(page)
if offset >= total or len(page) < _DELETE_SCAN_BATCH:
break
try:
total = int(reply[0])
except (TypeError, ValueError):
total = len(page)
if total <= len(page) or len(page) < _DELETE_SCAN_BATCH:
return deleted
return keys
raise RuntimeError(
f'Valkey deletion exceeded {_MAX_DELETE_SCAN_ROUNDS} batches ({_DELETE_SCAN_BATCH} keys per batch)'
)