diff --git a/CHANGELOG.md b/CHANGELOG.md index 8b57ec9a5..0d347cd53 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -143,6 +143,15 @@ construct the pipeline directly instead of subclassing the repository. No query behavior changes. +- **#1558**: Vector adapters are bound to the database, not to a project. `VectorIndexScope` + is the database namespace plus embedding schema; every write names the project it + touches (`upsert(project_id, ...)`, `delete(project_id, ...)`, + `delete_entity(project_id, ...)`, `delete_orphans(project_id, ...)`) and `search` takes + a `ProjectScope`, so one sqlite-vec or pgvector adapter answers a query across any set + of projects with one statement. Milvus keeps a collection per project and searches the + collections in scope. `SemanticSearch` passes its scope through, so a project + repository's vector search is unchanged. No query behavior changes. + ## v0.23.2 (2026-08-25) diff --git a/src/basic_memory/repository/milvus_index.py b/src/basic_memory/repository/milvus_index.py index 9cc367c15..dcc8e6ce2 100644 --- a/src/basic_memory/repository/milvus_index.py +++ b/src/basic_memory/repository/milvus_index.py @@ -13,6 +13,7 @@ MilvusStoredRecord, create_repository, ) +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_vector_index import ( VectorDeletion, VectorIndexScope, @@ -33,10 +34,10 @@ def _record_id(key: VectorKey) -> str: return hashlib.sha256(stable_key).hexdigest() -def collection_name(settings: MilvusSettings, scope: VectorIndexScope) -> str: - """Return a stable project collection name independent of embedding schema.""" +def collection_name(settings: MilvusSettings, scope: VectorIndexScope, project_id: int) -> str: + """Return one project's stable collection name, independent of embedding schema.""" namespace_digest = hashlib.sha256(scope.namespace.encode()).hexdigest()[:24] - return f"{settings.collection_prefix}_{namespace_digest}_{scope.project_id}" + return f"{settings.collection_prefix}_{namespace_digest}_{project_id}" def _normalize_cosine_score(score: float) -> float: @@ -45,7 +46,7 @@ def _normalize_cosine_score(score: float) -> float: class MilvusVectorIndex: - """Persist and query one Basic Memory project's vectors in Milvus.""" + """Persist and query Basic Memory vectors in Milvus, one collection per project.""" def __init__( self, @@ -56,10 +57,10 @@ def __init__( ) -> None: self.scope = scope self._settings = settings - self._collection_name = collection_name(settings, scope) self._repository_factory = repository_factory - self._initialized = False - self._initialize_lock = asyncio.Lock() + # Projects whose collection has been created or validated by this instance. + self._ready_projects: set[int] = set() + self._collection_lock = asyncio.Lock() def _with_repository[T](self, operation: Callable[[MilvusRepository], T]) -> T: repository = self._repository_factory(self._settings) @@ -68,27 +69,27 @@ def _with_repository[T](self, operation: Callable[[MilvusRepository], T]) -> T: finally: repository.close() - def _initialize_blocking(self) -> None: + def _collection(self, project_id: int) -> str: + return collection_name(self._settings, self.scope, project_id) + + def _validate_collection_blocking(self, collection: str) -> None: def initialize_repository(repository: MilvusRepository) -> None: - dimensions = repository.collection_dimensions(self._collection_name) + dimensions = repository.collection_dimensions(collection) if dimensions is None: - created = repository.create_collection( - self._collection_name, - self.scope.dimensions, - ) + created = repository.create_collection(collection, self.scope.dimensions) if created: return - dimensions = repository.collection_dimensions(self._collection_name) + dimensions = repository.collection_dimensions(collection) if dimensions is None: raise RuntimeError( - f"Milvus collection '{self._collection_name}' disappeared after " + f"Milvus collection '{collection}' disappeared after " "a concurrent create operation." ) if dimensions == self.scope.dimensions: # Milvus Lite releases persisted collections when the owning process exits. # Load only after the scope check so migrations do not load incompatible # remote collections before Basic Memory refuses to use them. - repository.load_collection(self._collection_name) + repository.load_collection(collection) return # Trigger: an existing project collection uses another embedding dimension. @@ -96,7 +97,7 @@ def initialize_repository(repository: MilvusRepository) -> None: # repeatedly erase each other's vectors during a rolling deployment. # Outcome: preserve the collection until an operator coordinates migration. raise RuntimeError( - f"Milvus collection '{self._collection_name}' has {dimensions} dimensions, " + f"Milvus collection '{collection}' has {dimensions} dimensions, " f"but Basic Memory is configured for {self.scope.dimensions}. Refusing to " "replace shared vector storage automatically; stop all writers and coordinate " "the collection migration before reindexing." @@ -106,12 +107,13 @@ def initialize_repository(repository: MilvusRepository) -> None: def _search_blocking( self, + collection: str, query: Sequence[float], limit: int, ) -> list[MilvusStoredMatch]: repository = self._repository_factory(self._settings) try: - return repository.search(self._collection_name, query, limit) + return repository.search(collection, query, limit) finally: repository.close() @@ -135,19 +137,28 @@ async def _run_blocking_mutation(self, operation: Callable[[], None]) -> None: raise async def initialize(self) -> None: - if self._initialized: - return - async with self._initialize_lock: - if self._initialized: - return - await self._run_blocking_mutation(self._initialize_blocking) - self._initialized = True + """Nothing is shared across projects: each collection is validated on first use.""" + return None + + async def _ensure_collection(self, project_id: int) -> str: + """Create or validate one project's collection once per adapter instance.""" + collection = self._collection(project_id) + if project_id in self._ready_projects: + return collection + async with self._collection_lock: + if project_id in self._ready_projects: + return collection + await self._run_blocking_mutation( + lambda: self._validate_collection_blocking(collection) + ) + self._ready_projects.add(project_id) + return collection - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: if not records: return validate_vector_dimensions(self.scope, records) - await self.initialize() + collection = await self._ensure_collection(project_id) stored_records = [ MilvusStoredRecord( record_id=_record_id(record.key), @@ -160,47 +171,44 @@ async def upsert(self, records: Sequence[VectorRecord]) -> None: ] await self._run_blocking_mutation( lambda: self._with_repository( - lambda repository: repository.upsert(self._collection_name, stored_records) + lambda repository: repository.upsert(collection, stored_records) ) ) - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: if not records: return - await self.initialize() + collection = await self._ensure_collection(project_id) stored_deletions = [(_record_id(record.key), record.source_hash) for record in records] await self._run_blocking_mutation( lambda: self._with_repository( - lambda repository: repository.delete_records( - self._collection_name, - stored_deletions, - ) + lambda repository: repository.delete_records(collection, stored_deletions) ) ) - async def delete_entity(self, entity_id: int) -> None: - await self.initialize() + async def delete_entity(self, project_id: int, entity_id: int) -> None: + collection = await self._ensure_collection(project_id) await self._run_blocking_mutation( lambda: self._with_repository( - lambda repository: repository.delete_entity(self._collection_name, entity_id) + lambda repository: repository.delete_entity(collection, entity_id) ) ) - async def delete_orphans(self, live_keys: Sequence[VectorKey]) -> None: - await self.initialize() + async def delete_orphans(self, project_id: int, live_keys: Sequence[VectorKey]) -> None: + collection = await self._ensure_collection(project_id) live_ids = {_record_id(key) for key in live_keys} def delete_missing(repository: MilvusRepository) -> None: orphan_ids: list[str] = [] - for record_id in repository.iter_ids(self._collection_name): + for record_id in repository.iter_ids(collection): if record_id in live_ids: continue orphan_ids.append(record_id) if len(orphan_ids) == _ORPHAN_DELETE_BATCH_SIZE: - repository.delete_ids(self._collection_name, orphan_ids) + repository.delete_ids(collection, orphan_ids) orphan_ids.clear() if orphan_ids: - repository.delete_ids(self._collection_name, orphan_ids) + repository.delete_ids(collection, orphan_ids) await self._run_blocking_mutation(lambda: self._with_repository(delete_missing)) @@ -209,28 +217,32 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: - if not query or limit <= 0: + if not query or limit <= 0 or projects.is_empty: return [] validate_query_dimensions(self.scope, query) - await self.initialize() - - stored_matches = await asyncio.to_thread(self._search_blocking, query, limit) - matches = [ - VectorMatch( - key=VectorKey( - entity_id=match.entity_id, - chunk_key=match.chunk_key, - ), - similarity=_normalize_cosine_score(match.score), + + # Milvus has no cross-collection search, so a scope wider than one project + # asks each project's collection for its own top ``limit`` and merges them. + matches: list[VectorMatch] = [] + for project_id in projects.project_ids: + collection = await self._ensure_collection(project_id) + stored_matches = await asyncio.to_thread( + self._search_blocking, collection, query, limit ) - for match in stored_matches - ] - return sorted( - matches, + matches.extend( + VectorMatch( + key=VectorKey(entity_id=match.entity_id, chunk_key=match.chunk_key), + similarity=_normalize_cosine_score(match.score), + ) + for match in stored_matches + ) + matches.sort( key=lambda match: ( -match.similarity, match.key.entity_id, match.key.chunk_key, - ), + ) ) + return matches[:limit] diff --git a/src/basic_memory/repository/pgvector_index.py b/src/basic_memory/repository/pgvector_index.py index 764a1c716..1db0be8c5 100644 --- a/src/basic_memory/repository/pgvector_index.py +++ b/src/basic_memory/repository/pgvector_index.py @@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from basic_memory import db +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError from basic_memory.repository.semantic_vector_index import ( VectorDeletion, @@ -151,37 +152,7 @@ async def _has_source_hash_column(self, session: AsyncSession) -> bool: ) return result.scalar_one_or_none() is not None - async def _chunk_ids_by_key( - self, - session: AsyncSession, - keys: Sequence[VectorKey], - ) -> dict[VectorKey, int]: - if not keys: - return {} - - params: dict[str, object] = {"project_id": self.scope.project_id} - predicates: list[str] = [] - for index, key in enumerate(keys): - params[f"entity_id_{index}"] = key.entity_id - params[f"chunk_key_{index}"] = key.chunk_key - predicates.append( - f"(entity_id = :entity_id_{index} AND chunk_key = :chunk_key_{index})" - ) - result = await session.execute( - text( - "SELECT id, entity_id, chunk_key FROM search_vector_chunks " - "WHERE project_id = :project_id AND (" + " OR ".join(predicates) + ")" - ), - params, - ) - return { - VectorKey(entity_id=int(row["entity_id"]), chunk_key=str(row["chunk_key"])): int( - row["id"] - ) - for row in result.mappings().all() - } - - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: if not records: return validate_vector_dimensions(self.scope, records) @@ -189,7 +160,7 @@ async def upsert(self, records: Sequence[VectorRecord]) -> None: async with db.scoped_session(self._session_maker) as session: keys = [record.key for record in records] - params: dict[str, object] = {"project_id": self.scope.project_id} + params: dict[str, object] = {"project_id": project_id} predicates: list[str] = [] for index, key in enumerate(keys): params[f"entity_id_{index}"] = key.entity_id @@ -224,7 +195,7 @@ async def upsert(self, records: Sequence[VectorRecord]) -> None: if not current_records: return - params = {"project_id": self.scope.project_id} + params = {"project_id": project_id} values: list[str] = [] for index, record in enumerate(current_records): params[f"chunk_id_{index}"] = manifest_by_key[record.key][0] @@ -252,12 +223,12 @@ async def upsert(self, records: Sequence[VectorRecord]) -> None: ) await session.commit() - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: if not records: return await self.initialize() async with db.scoped_session(self._session_maker) as session: - params: dict[str, object] = {"project_id": self.scope.project_id} + params: dict[str, object] = {"project_id": project_id} predicates: list[str] = [] for index, record in enumerate(records): params[f"entity_id_{index}"] = record.key.entity_id @@ -294,7 +265,7 @@ async def delete(self, records: Sequence[VectorDeletion]) -> None: ) await session.commit() - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: await self.initialize() async with db.scoped_session(self._session_maker) as session: await session.execute( @@ -303,12 +274,12 @@ async def delete_entity(self, entity_id: int) -> None: "SELECT id FROM search_vector_chunks " "WHERE project_id = :project_id AND entity_id = :entity_id)" ), - {"project_id": self.scope.project_id, "entity_id": entity_id}, + {"project_id": project_id, "entity_id": entity_id}, ) await session.commit() - async def delete_orphans(self, _live_keys: Sequence[VectorKey]) -> None: - """Remove pgvector rows absent from the current ready manifest scope.""" + async def delete_orphans(self, project_id: int, _live_keys: Sequence[VectorKey]) -> None: + """Remove pgvector rows absent from one project's current ready manifest.""" await self.initialize() async with db.scoped_session(self._session_maker) as session: await session.execute( @@ -324,7 +295,7 @@ async def delete_orphans(self, _live_keys: Sequence[VectorKey]) -> None: "AND chunks.embedding_status = 'ready')" ), { - "project_id": self.scope.project_id, + "project_id": project_id, "embedding_identity": self.scope.embedding_identity, }, ) @@ -335,11 +306,21 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: - if not query or limit <= 0: + if not query or limit <= 0 or projects.is_empty: return [] validate_query_dimensions(self.scope, query) await self.initialize() + params: dict[str, object] = { + "query": self._format_vector(query), + "dimensions": self.scope.dimensions, + "embedding_identity": self.scope.embedding_identity, + "limit": limit, + } + # The scope binds its ids once; both predicates reference the same names. + embeddings_in_scope = projects.predicate("e.project_id", params) + chunks_in_scope = projects.predicate("c.project_id", params) async with db.scoped_session(self._session_maker) as session: result = await session.execute( text( @@ -347,9 +328,9 @@ async def search( "1 - (e.embedding <=> CAST(:query AS vector)) AS similarity " "FROM search_vector_embeddings e " "JOIN search_vector_chunks c ON c.id = e.chunk_id " - "WHERE e.project_id = :project_id " + f"WHERE {embeddings_in_scope} " "AND e.embedding_dims = :dimensions " - "AND c.project_id = :project_id " + f"AND {chunks_in_scope} " "AND c.vector_index = 'pgvector' " "AND c.embedding_status = 'ready' " "AND c.embedding_model = :embedding_identity " @@ -358,13 +339,7 @@ async def search( "c.entity_id ASC, c.chunk_key ASC " "LIMIT :limit" ), - { - "query": self._format_vector(query), - "project_id": self.scope.project_id, - "dimensions": self.scope.dimensions, - "embedding_identity": self.scope.embedding_identity, - "limit": limit, - }, + params, ) return [ VectorMatch( diff --git a/src/basic_memory/repository/postgres_search_repository.py b/src/basic_memory/repository/postgres_search_repository.py index 06f4e57d8..b57b4eb01 100644 --- a/src/basic_memory/repository/postgres_search_repository.py +++ b/src/basic_memory/repository/postgres_search_repository.py @@ -113,11 +113,7 @@ def __init__( ) vector_index = PgVectorIndex( session_maker, - build_vector_index_scope( - self._app_config, - self._embedding_provider, - project_id, - ), + build_vector_index_scope(self._app_config, self._embedding_provider), ) self._semantic_vector_index_name = effective_name self._semantic_vector_index = vector_index @@ -314,11 +310,7 @@ async def _ensure_vector_tables(self) -> None: self._semantic_vector_index_name = "pgvector" self._semantic_vector_index = PgVectorIndex( self.session_maker, - build_vector_index_scope( - self._app_config, - self._embedding_provider, - self.project_id, - ), + build_vector_index_scope(self._app_config, self._embedding_provider), ) if self._vector_tables_initialized: return diff --git a/src/basic_memory/repository/search_reader.py b/src/basic_memory/repository/search_reader.py index d5f925d54..0b5c81eb5 100644 --- a/src/basic_memory/repository/search_reader.py +++ b/src/basic_memory/repository/search_reader.py @@ -260,7 +260,9 @@ async def _run_vector_query( return [] if not self.vector.external: - matches = await self.vector.index.search(query_embedding, limit=candidate_limit) + matches = await self.vector.index.search( + query_embedding, limit=candidate_limit, projects=self.scope + ) if trace is not None: trace.readiness = await read_manifest_readiness( session, @@ -272,7 +274,9 @@ async def _run_vector_query( scan_limit = min(candidate_limit, VECTOR_FILTER_SCAN_LIMIT) while True: - matches = await self.vector.index.search(query_embedding, limit=scan_limit) + matches = await self.vector.index.search( + query_embedding, limit=scan_limit, projects=self.scope + ) if trace is not None and trace.readiness is None: trace.readiness = await read_manifest_readiness( session, diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index 2b9828202..b21aca5b1 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -235,7 +235,6 @@ def create_search_repository( embedding_provider = create_embedding_provider(config) vector_index_name, vector_index = create_semantic_vector_index( session_maker=session_maker, - project_id=project_id, app_config=config, database_backend=database_backend, embedding_provider=embedding_provider, diff --git a/src/basic_memory/repository/search_repository_base.py b/src/basic_memory/repository/search_repository_base.py index 15e2dd506..8935b307d 100644 --- a/src/basic_memory/repository/search_repository_base.py +++ b/src/basic_memory/repository/search_repository_base.py @@ -621,7 +621,7 @@ async def _persist_embeddings( # the authoritative SQL database. Hold its manifest lock across # adapter I/O so a newer prepare cannot advance this generation # before the external write and ready transition complete. - await self._semantic_vector_index.upsert(records) + await self._semantic_vector_index.upsert(self.project_id, records) await self._mark_embedding_jobs_ready( session, params=params, @@ -632,7 +632,7 @@ async def _persist_embeddings( # Built-in adapters share the authoritative database. They verify and lock # each record's source_hash inside the same transaction as their vector write. - await self._semantic_vector_index.upsert(records) + await self._semantic_vector_index.upsert(self.project_id, records) async with db.scoped_session(self.session_maker) as session: await self._mark_embedding_jobs_ready( session, @@ -834,7 +834,7 @@ async def _finalize_prepared_vector_deletions( ] if not current_deletions: return - await self._semantic_vector_index.delete(current_deletions) + await self._semantic_vector_index.delete(self.project_id, current_deletions) await session.execute( text( "DELETE FROM search_vector_chunks " @@ -846,7 +846,7 @@ async def _finalize_prepared_vector_deletions( await session.commit() return - await self._semantic_vector_index.delete(deletions) + await self._semantic_vector_index.delete(self.project_id, deletions) if self._semantic_vector_index_name in BUILT_IN_VECTOR_INDEX_NAMES: return async with db.scoped_session(self.session_maker) as session: @@ -1209,7 +1209,7 @@ async def _delete_external_entity_vectors_locked( await self._semantic_vector_index.initialize() for entity_id in deleted_entity_ids: - await self._semantic_vector_index.delete_entity(entity_id) + await self._semantic_vector_index.delete_entity(self.project_id, entity_id) async def _delete_project_builtin_vector_rows(self, session: AsyncSession) -> None: """Delete backend-owned vector rows before their SQL manifest is removed.""" @@ -1479,11 +1479,11 @@ async def reconcile_vector_index(self) -> None: ] if external_vector_index: - await self._semantic_vector_index.delete_orphans(live_keys) + await self._semantic_vector_index.delete_orphans(self.project_id, live_keys) await session.commit() return - await self._semantic_vector_index.delete_orphans(live_keys) + await self._semantic_vector_index.delete_orphans(self.project_id, live_keys) # ------------------------------------------------------------------ # Shared semantic search: guard, text processing, chunking diff --git a/src/basic_memory/repository/semantic_vector_index.py b/src/basic_memory/repository/semantic_vector_index.py index e869fb0a8..bb6b5a2b7 100644 --- a/src/basic_memory/repository/semantic_vector_index.py +++ b/src/basic_memory/repository/semantic_vector_index.py @@ -6,25 +6,28 @@ from dataclasses import dataclass from typing import Protocol, runtime_checkable +from basic_memory.repository.search_scope import ProjectScope + @dataclass(frozen=True, slots=True) class VectorIndexScope: - """Stable project storage identity plus the current embedding schema.""" + """The database namespace and embedding schema an adapter stores vectors under. + + Projects are partitions inside it: every write names the project it touches and a + search names the projects it reads, so one adapter serves a whole database. + """ namespace: str - project_id: int embedding_identity: str dimensions: int - @property - def storage_key(self) -> tuple[str, int]: - """Return the stable isolation key external adapters must use for storage.""" - return (self.namespace, self.project_id) - @dataclass(frozen=True, slots=True) class VectorKey: - """Backend-independent identity for one semantic chunk vector.""" + """Backend-independent identity for one semantic chunk vector. + + Entity ids are database-wide primary keys, so the pair is unique across projects. + """ entity_id: int chunk_key: str @@ -68,19 +71,19 @@ class SemanticVectorIndex(Protocol): def scope(self) -> VectorIndexScope: ... async def initialize(self) -> None: - """Create or validate backend storage for the configured scope.""" + """Create or validate backend storage shared by every project in the scope.""" ... - async def upsert(self, records: Sequence[VectorRecord]) -> None: - """Insert or replace vectors only for each record's source generation.""" + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: + """Insert or replace one project's vectors, only for each record's source generation.""" ... - async def delete(self, records: Sequence[VectorDeletion]) -> None: - """Delete vectors only for each record's source generation.""" + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: + """Delete one project's vectors, only for each record's source generation.""" ... - async def delete_entity(self, entity_id: int) -> None: - """Delete every vector owned by an entity in this scope.""" + async def delete_entity(self, project_id: int, entity_id: int) -> None: + """Delete every vector owned by an entity in one project.""" ... async def search( @@ -88,8 +91,9 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: - """Return nearest matches ordered by normalized cosine similarity.""" + """Return nearest matches within ``projects``, ordered by normalized similarity.""" ... @@ -100,8 +104,8 @@ class SemanticVectorIndexReconciler(Protocol): @property def scope(self) -> VectorIndexScope: ... - async def delete_orphans(self, live_keys: Sequence[VectorKey]) -> None: - """Delete scoped vectors whose stable keys are not in ``live_keys``.""" + async def delete_orphans(self, project_id: int, live_keys: Sequence[VectorKey]) -> None: + """Delete one project's vectors whose stable keys are not in ``live_keys``.""" ... diff --git a/src/basic_memory/repository/semantic_vector_index_factory.py b/src/basic_memory/repository/semantic_vector_index_factory.py index d9bd91869..f29d00517 100644 --- a/src/basic_memory/repository/semantic_vector_index_factory.py +++ b/src/basic_memory/repository/semantic_vector_index_factory.py @@ -76,12 +76,13 @@ def _database_namespace(app_config: BasicMemoryConfig) -> str: def build_vector_index_scope( app_config: BasicMemoryConfig, provider: EmbeddingProvider, - project_id: int, ) -> VectorIndexScope: - """Build the explicit isolation contract handed to every vector adapter.""" + """Build the storage identity handed to every vector adapter. + + The database namespace and embedding schema; projects are named per operation. + """ return VectorIndexScope( namespace=_database_namespace(app_config), - project_id=project_id, embedding_identity=semantic_embedding_identity(provider), dimensions=provider.dimensions, ) @@ -109,14 +110,13 @@ def _create_milvus_index( def create_semantic_vector_index( *, session_maker: async_sessionmaker[AsyncSession], - project_id: int, app_config: BasicMemoryConfig, database_backend: DatabaseBackend, embedding_provider: EmbeddingProvider, ) -> tuple[str, SemanticVectorIndex]: """Create the vector adapter selected by the validated application config.""" name = resolve_semantic_vector_index_name(app_config, database_backend) - scope = build_vector_index_scope(app_config, embedding_provider, project_id) + scope = build_vector_index_scope(app_config, embedding_provider) if name == "sqlite-vec": from basic_memory.repository.sqlite_vec_index import SQLiteVecIndex diff --git a/src/basic_memory/repository/sqlite_search_repository.py b/src/basic_memory/repository/sqlite_search_repository.py index 0861023b7..7f50b8df5 100644 --- a/src/basic_memory/repository/sqlite_search_repository.py +++ b/src/basic_memory/repository/sqlite_search_repository.py @@ -85,11 +85,7 @@ def __init__( self._vector_dimensions = self._embedding_provider.dimensions self._semantic_vector_index = vector_index or SQLiteVecIndex( session_maker, - build_vector_index_scope( - self._app_config, - self._embedding_provider, - project_id, - ), + build_vector_index_scope(self._app_config, self._embedding_provider), ) @override @@ -269,11 +265,7 @@ async def _ensure_vector_tables(self) -> None: assert self._embedding_provider is not None self._semantic_vector_index = SQLiteVecIndex( self.session_maker, - build_vector_index_scope( - self._app_config, - self._embedding_provider, - self.project_id, - ), + build_vector_index_scope(self._app_config, self._embedding_provider), ) if self._vector_tables_initialized: return diff --git a/src/basic_memory/repository/sqlite_vec_index.py b/src/basic_memory/repository/sqlite_vec_index.py index c3344c31c..f85f40248 100644 --- a/src/basic_memory/repository/sqlite_vec_index.py +++ b/src/basic_memory/repository/sqlite_vec_index.py @@ -13,6 +13,7 @@ from basic_memory import db from basic_memory.models.search import create_sqlite_search_vector_embeddings +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError from basic_memory.repository.semantic_vector_index import ( VectorDeletion, @@ -126,43 +127,14 @@ async def initialize(self) -> None: await session.commit() self._initialized = True - async def _rowids_by_key( - self, - session: AsyncSession, - keys: Sequence[VectorKey], - ) -> dict[VectorKey, int]: - if not keys: - return {} - params: dict[str, object] = {"project_id": self.scope.project_id} - predicates: list[str] = [] - for index, key in enumerate(keys): - params[f"entity_id_{index}"] = key.entity_id - params[f"chunk_key_{index}"] = key.chunk_key - predicates.append( - f"(entity_id = :entity_id_{index} AND chunk_key = :chunk_key_{index})" - ) - result = await session.execute( - text( - "SELECT id, entity_id, chunk_key FROM search_vector_chunks " - "WHERE project_id = :project_id AND (" + " OR ".join(predicates) + ")" - ), - params, - ) - return { - VectorKey(entity_id=int(row["entity_id"]), chunk_key=str(row["chunk_key"])): int( - row["id"] - ) - for row in result.mappings().all() - } - - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: if not records: return validate_vector_dimensions(self.scope, records) await self.initialize() async with db.scoped_session(self._session_maker) as session: await self._ensure_loaded(session) - params: dict[str, object] = {"project_id": self.scope.project_id} + params: dict[str, object] = {"project_id": project_id} predicates: list[str] = [] records_by_key = {record.key: record for record in records} for index, record in enumerate(records): @@ -223,13 +195,13 @@ async def upsert(self, records: Sequence[VectorRecord]) -> None: ) await session.commit() - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: if not records: return await self.initialize() async with db.scoped_session(self._session_maker) as session: await self._ensure_loaded(session) - params: dict[str, object] = {"project_id": self.scope.project_id} + params: dict[str, object] = {"project_id": project_id} predicates: list[str] = [] for index, record in enumerate(records): params[f"entity_id_{index}"] = record.key.entity_id @@ -264,7 +236,7 @@ async def delete(self, records: Sequence[VectorDeletion]) -> None: ) await session.commit() - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: await self.initialize() async with db.scoped_session(self._session_maker) as session: await self._ensure_loaded(session) @@ -274,12 +246,12 @@ async def delete_entity(self, entity_id: int) -> None: "SELECT id FROM search_vector_chunks " "WHERE project_id = :project_id AND entity_id = :entity_id)" ), - {"project_id": self.scope.project_id, "entity_id": entity_id}, + {"project_id": project_id, "entity_id": entity_id}, ) await session.commit() - async def delete_orphans(self, _live_keys: Sequence[VectorKey]) -> None: - """Remove sqlite-vec rows absent from the current ready manifest scope.""" + async def delete_orphans(self, project_id: int, _live_keys: Sequence[VectorKey]) -> None: + """Remove sqlite-vec rows absent from one project's current ready manifest.""" await self.initialize() async with db.scoped_session(self._session_maker) as session: await self._ensure_loaded(session) @@ -316,7 +288,7 @@ async def delete_orphans(self, _live_keys: Sequence[VectorKey]) -> None: "AND embedding_status = 'ready'))" ), { - "project_id": self.scope.project_id, + "project_id": project_id, "embedding_identity": self.scope.embedding_identity, }, ) @@ -327,12 +299,20 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: - if not query or limit <= 0: + if not query or limit <= 0 or projects.is_empty: return [] validate_query_dimensions(self.scope, query) await self.initialize() vector_k = min(limit, SQLITE_VEC_MAX_K) + params: dict[str, object] = { + "query": json.dumps(list(query)), + "vector_k": vector_k, + "embedding_identity": self.scope.embedding_identity, + "limit": limit, + } + chunks_in_scope = projects.predicate("c.project_id", params) async with db.scoped_session(self._session_maker) as session: await self._ensure_loaded(session) result = await session.execute( @@ -345,20 +325,14 @@ async def search( "FROM vector_matches " "JOIN search_vector_chunks c ON c.id = vector_matches.rowid " "AND c.source_hash = vector_matches.source_hash " - "WHERE c.project_id = :project_id " + f"WHERE {chunks_in_scope} " "AND c.vector_index = 'sqlite-vec' " "AND c.embedding_status = 'ready' " "AND c.embedding_model = :embedding_identity " "ORDER BY vector_matches.distance ASC, " "c.entity_id ASC, c.chunk_key ASC LIMIT :limit" ), - { - "query": json.dumps(list(query)), - "vector_k": vector_k, - "project_id": self.scope.project_id, - "embedding_identity": self.scope.embedding_identity, - "limit": limit, - }, + params, ) return [ VectorMatch( diff --git a/test-int/semantic/conftest.py b/test-int/semantic/conftest.py index c5c87559a..53ea89946 100644 --- a/test-int/semantic/conftest.py +++ b/test-int/semantic/conftest.py @@ -293,7 +293,6 @@ async def create_search_service( if embedding_provider is not None: vector_index_name, vector_index = create_semantic_vector_index( session_maker=session_maker, - project_id=project.id, app_config=app_config, database_backend=combo.backend, embedding_provider=embedding_provider, diff --git a/test-int/semantic/test_milvus_lite.py b/test-int/semantic/test_milvus_lite.py index 7f4db5e08..a637204c6 100644 --- a/test-int/semantic/test_milvus_lite.py +++ b/test-int/semantic/test_milvus_lite.py @@ -11,6 +11,7 @@ from basic_memory.repository.milvus_config import MilvusSettings from basic_memory.repository.milvus_index import MilvusVectorIndex +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_vector_index import ( VectorDeletion, VectorIndexScope, @@ -23,6 +24,9 @@ pytest.mark.skipif(sys.platform == "win32", reason="Milvus Lite does not support Windows"), ] +PROJECT = 7 +PROJECTS = ProjectScope.single(PROJECT) + _RESTART_SCRIPT = """ import asyncio @@ -30,6 +34,7 @@ from basic_memory.repository.milvus_config import MilvusSettings from basic_memory.repository.milvus_index import MilvusVectorIndex +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_vector_index import ( VectorIndexScope, VectorKey, @@ -41,7 +46,6 @@ async def main() -> None: phase, database_path = sys.argv[1:] scope = VectorIndexScope( namespace="restart-database", - project_id=7, embedding_identity="Stub:model", dimensions=3, ) @@ -50,18 +54,19 @@ async def main() -> None: if phase == "write": await index.upsert( + 7, [ VectorRecord( key=key, source_hash="auth-v1", values=(1.0, 0.0, 0.0), ) - ] + ], ) return - await index.delete_orphans([key]) - matches = await index.search((1.0, 0.0, 0.0), limit=1) + await index.delete_orphans(7, [key]) + matches = await index.search((1.0, 0.0, 0.0), limit=1, projects=ProjectScope.single(7)) assert [match.key for match in matches] == [key] @@ -83,7 +88,6 @@ def _run_restart_phase(phase: str, database_path: str) -> None: async def test_milvus_lite_vector_lifecycle(tmp_path) -> None: scope = VectorIndexScope( namespace="integration-database", - project_id=7, embedding_identity="Stub:model", dimensions=3, ) @@ -94,6 +98,7 @@ async def test_milvus_lite_vector_lifecycle(tmp_path) -> None: auth_key = VectorKey(entity_id=1, chunk_key="summary:0") database_key = VectorKey(entity_id=2, chunk_key="summary:0") await index.upsert( + PROJECT, [ VectorRecord( key=auth_key, @@ -105,21 +110,23 @@ async def test_milvus_lite_vector_lifecycle(tmp_path) -> None: source_hash="database-v1", values=(0.0, 1.0, 0.0), ), - ] + ], ) - matches = await index.search((1.0, 0.0, 0.0), limit=2) + matches = await index.search((1.0, 0.0, 0.0), limit=2, projects=PROJECTS) assert matches[0].key == auth_key assert matches[0].similarity == pytest.approx(1.0) - await index.delete([VectorDeletion(key=auth_key, source_hash="stale-generation")]) - assert (await index.search((1.0, 0.0, 0.0), limit=2))[0].key == auth_key + await index.delete(PROJECT, [VectorDeletion(key=auth_key, source_hash="stale-generation")]) + assert (await index.search((1.0, 0.0, 0.0), limit=2, projects=PROJECTS))[0].key == auth_key - await index.delete([VectorDeletion(key=auth_key, source_hash="auth-v1")]) - assert [match.key for match in await index.search((1.0, 0.0, 0.0), limit=2)] == [database_key] + await index.delete(PROJECT, [VectorDeletion(key=auth_key, source_hash="auth-v1")]) + assert [ + match.key for match in await index.search((1.0, 0.0, 0.0), limit=2, projects=PROJECTS) + ] == [database_key] - await index.delete_orphans([]) - assert await index.search((1.0, 0.0, 0.0), limit=2) == [] + await index.delete_orphans(PROJECT, []) + assert await index.search((1.0, 0.0, 0.0), limit=2, projects=PROJECTS) == [] def test_milvus_lite_reloads_collection_after_process_restart(tmp_path) -> None: diff --git a/test-int/test_embedding_status_vec0.py b/test-int/test_embedding_status_vec0.py index 1ccbee5cc..d4c3ee949 100644 --- a/test-int/test_embedding_status_vec0.py +++ b/test-int/test_embedding_status_vec0.py @@ -137,13 +137,14 @@ async def test_embedding_status_reads_real_vec0_table(engine_factory, test_proje # An obsolete embedding result must not claim the stable vec0 row after the # manifest has advanced to a newer source generation. await search_repo._semantic_vector_index.upsert( + project_id, [ VectorRecord( key=VectorKey(entity_id=entity_id, chunk_key="chunk-1"), source_hash="stale-hash", values=tuple(_unit_vector(dimensions)), ) - ] + ], ) async with db.scoped_session(session_maker) as session: # sqlite-vec is loaded per connection. Windows may hand this assertion a @@ -153,13 +154,14 @@ async def test_embedding_status_reads_real_vec0_table(engine_factory, test_proje assert stale_count.scalar_one() == 0 await search_repo._semantic_vector_index.upsert( + project_id, [ VectorRecord( key=VectorKey(entity_id=entity_id, chunk_key="chunk-1"), source_hash="hash", values=tuple(_unit_vector(dimensions)), ) - ] + ], ) async with db.scoped_session(session_maker) as session: diff --git a/tests/repository/test_milvus_index.py b/tests/repository/test_milvus_index.py index 47790c72d..0dfe584c5 100644 --- a/tests/repository/test_milvus_index.py +++ b/tests/repository/test_milvus_index.py @@ -23,6 +23,7 @@ MilvusVectorIndex, collection_name, ) +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_vector_index import ( SemanticVectorIndex, SemanticVectorIndexReconciler, @@ -35,6 +36,11 @@ create_semantic_vector_index, ) +PROJECT = 42 +OTHER_PROJECT = 43 +PROJECTS = ProjectScope.single(PROJECT) +QUERY = [1.0, 0.0, 0.0] + class FakeRepository: """In-memory recorder for the blocking Milvus repository.""" @@ -57,6 +63,8 @@ def __init__( self.ids: list[str] = [] self.id_deletes: list[tuple[str, list[str]]] = [] self.matches: list[MilvusStoredMatch] = [] + # Per-collection answers; ``matches`` is the answer for any collection not listed. + self.matches_by_collection: dict[str, list[MilvusStoredMatch]] = {} self.searches: list[tuple[str, list[float], int]] = [] self.closed = 0 @@ -103,7 +111,7 @@ def search( limit: int, ) -> list[MilvusStoredMatch]: self.searches.append((collection_name, list(query), limit)) - return self.matches + return self.matches_by_collection.get(collection_name, self.matches) def close(self) -> None: self.closed += 1 @@ -172,7 +180,6 @@ def runtime_log_attrs(self) -> dict[str, Any]: def scope() -> VectorIndexScope: return VectorIndexScope( namespace="basic-memory-database", - project_id=42, embedding_identity="Provider:model-a", dimensions=3, ) @@ -195,84 +202,99 @@ def _index( ) -def test_collection_name_uses_only_stable_scope_identity( +def test_collection_name_uses_only_stable_project_identity( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: changed_schema = VectorIndexScope( namespace=scope.namespace, - project_id=scope.project_id, embedding_identity="Provider:model-b", dimensions=9, ) - other_project = VectorIndexScope( - namespace=scope.namespace, - project_id=43, - embedding_identity=scope.embedding_identity, - dimensions=scope.dimensions, + + assert collection_name(settings, scope, PROJECT) == collection_name( + settings, changed_schema, PROJECT ) + assert collection_name(settings, scope, PROJECT) != collection_name( + settings, scope, OTHER_PROJECT + ) + assert collection_name(settings, scope, PROJECT).startswith("basic_memory_") + + +# --- Collection validation happens on a project's first use --- + + +@pytest.mark.asyncio +async def test_initialize_prepares_nothing_shared( + scope: VectorIndexScope, + settings: MilvusSettings, +) -> None: + """Collections are per project, so the database-wide hook has nothing to create.""" + repository = FakeRepository() - assert collection_name(settings, scope) == collection_name(settings, changed_schema) - assert collection_name(settings, scope) != collection_name(settings, other_project) - assert collection_name(settings, scope).startswith("basic_memory_") + await _index(scope, settings, repository).initialize() + + assert repository.created == [] + assert repository.closed == 0 @pytest.mark.asyncio -async def test_initialize_creates_missing_collection_once( +async def test_first_use_creates_missing_collection_once( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository() index = _index(scope, settings, repository) - await index.initialize() - await index.initialize() + await index.search(QUERY, limit=1, projects=PROJECTS) + await index.search(QUERY, limit=1, projects=PROJECTS) - assert repository.created == [(collection_name(settings, scope), scope.dimensions)] - assert repository.closed == 1 + assert repository.created == [(collection_name(settings, scope, PROJECT), scope.dimensions)] + # One validation plus one search per call. + assert repository.closed == 3 @pytest.mark.asyncio -async def test_initialize_accepts_compatible_collection_create_race( +async def test_first_use_accepts_compatible_collection_create_race( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository(create_result=False, race_dimensions=scope.dimensions) - await _index(scope, settings, repository).initialize() + await _index(scope, settings, repository).search(QUERY, limit=1, projects=PROJECTS) - assert repository.created == [(collection_name(settings, scope), scope.dimensions)] + assert repository.created == [(collection_name(settings, scope, PROJECT), scope.dimensions)] assert repository.dimensions == scope.dimensions - assert repository.loaded == [collection_name(settings, scope)] + assert repository.loaded == [collection_name(settings, scope, PROJECT)] @pytest.mark.asyncio -async def test_initialize_rejects_incompatible_collection_create_race( +async def test_first_use_rejects_incompatible_collection_create_race( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository(create_result=False, race_dimensions=99) with pytest.raises(RuntimeError, match="Refusing to replace shared vector storage"): - await _index(scope, settings, repository).initialize() + await _index(scope, settings, repository).search(QUERY, limit=1, projects=PROJECTS) assert repository.dimensions == 99 assert repository.loaded == [] @pytest.mark.asyncio -async def test_initialize_rejects_collection_disappearing_after_create_race( +async def test_first_use_rejects_collection_disappearing_after_create_race( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository(create_result=False) with pytest.raises(RuntimeError, match="disappeared after a concurrent create"): - await _index(scope, settings, repository).initialize() + await _index(scope, settings, repository).search(QUERY, limit=1, projects=PROJECTS) @pytest.mark.asyncio -async def test_initialize_preserves_collection_on_dimension_mismatch( +async def test_first_use_preserves_collection_on_dimension_mismatch( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: @@ -280,7 +302,7 @@ async def test_initialize_preserves_collection_on_dimension_mismatch( index = _index(scope, settings, repository) with pytest.raises(RuntimeError, match="Refusing to replace shared vector storage"): - await index.initialize() + await index.search(QUERY, limit=1, projects=PROJECTS) assert repository.created == [] assert repository.dimensions == 99 @@ -288,36 +310,39 @@ async def test_initialize_preserves_collection_on_dimension_mismatch( @pytest.mark.asyncio -async def test_initialize_accepts_matching_collection( +async def test_first_use_accepts_matching_collection( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository(dimensions=scope.dimensions) - await _index(scope, settings, repository).initialize() + await _index(scope, settings, repository).search(QUERY, limit=1, projects=PROJECTS) assert repository.created == [] - assert repository.loaded == [collection_name(settings, scope)] + assert repository.loaded == [collection_name(settings, scope, PROJECT)] @pytest.mark.asyncio -async def test_concurrent_initialize_rechecks_state_inside_lock( +async def test_concurrent_first_use_rechecks_state_inside_lock( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository() index = _index(scope, settings, repository) - await index._initialize_lock.acquire() - waiting_initialize = asyncio.create_task(index.initialize()) + await index._collection_lock.acquire() + waiting = asyncio.create_task(index._ensure_collection(PROJECT)) await asyncio.sleep(0) - index._initialized = True - index._initialize_lock.release() - await waiting_initialize + index._ready_projects.add(PROJECT) + index._collection_lock.release() + assert await waiting == collection_name(settings, scope, PROJECT) assert repository.closed == 0 +# --- Writes --- + + @pytest.mark.asyncio async def test_upsert_preserves_stable_key_generation_and_values( scope: VectorIndexScope, @@ -331,9 +356,10 @@ async def test_upsert_preserves_stable_key_generation_and_values( values=(1.0, 0.0, -1.0), ) - await index.upsert([record]) + await index.upsert(PROJECT, [record]) - _, stored_records = repository.upserts[0] + collection, stored_records = repository.upserts[0] + assert collection == collection_name(settings, scope, PROJECT) assert len(stored_records) == 1 assert stored_records[0].entity_id == 7 assert stored_records[0].chunk_key == "summary:0" @@ -352,13 +378,14 @@ async def test_upsert_rejects_wrong_dimensions_before_milvus_call( with pytest.raises(ValueError, match="expected 3, got 2"): await index.upsert( + PROJECT, [ VectorRecord( key=VectorKey(entity_id=7, chunk_key="summary:0"), source_hash="source-a", values=(1.0, 0.0), ) - ] + ], ) assert repository.upserts == [] @@ -378,14 +405,14 @@ async def test_mutations_finish_before_propagating_cancellation( if operation == "upsert": mutation = index.upsert( - [VectorRecord(key=key, source_hash="source-a", values=(1.0, 0.0, 0.0))] + PROJECT, [VectorRecord(key=key, source_hash="source-a", values=(1.0, 0.0, 0.0))] ) elif operation == "delete": - mutation = index.delete([VectorDeletion(key=key, source_hash="source-a")]) + mutation = index.delete(PROJECT, [VectorDeletion(key=key, source_hash="source-a")]) elif operation == "delete_entity": - mutation = index.delete_entity(key.entity_id) + mutation = index.delete_entity(PROJECT, key.entity_id) else: - mutation = index.delete_orphans([]) + mutation = index.delete_orphans(PROJECT, []) mutation_task = asyncio.create_task(mutation) async with asyncio.timeout(2): @@ -415,24 +442,29 @@ async def test_delete_forwards_source_generation( index = _index(scope, settings, repository) key = VectorKey(entity_id=7, chunk_key="summary:0") - await index.delete([VectorDeletion(key=key, source_hash="source-a")]) + await index.delete(PROJECT, [VectorDeletion(key=key, source_hash="source-a")]) - _, deletions = repository.record_deletes[0] + collection, deletions = repository.record_deletes[0] + assert collection == collection_name(settings, scope, PROJECT) assert deletions[0][1] == "source-a" assert len(deletions[0][0]) == 64 @pytest.mark.asyncio -async def test_delete_entity_uses_project_collection( +async def test_delete_entity_uses_the_named_project_collection( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository(dimensions=scope.dimensions) index = _index(scope, settings, repository) - await index.delete_entity(77) + await index.delete_entity(PROJECT, 77) + await index.delete_entity(OTHER_PROJECT, 78) - assert repository.entity_deletes == [(collection_name(settings, scope), 77)] + assert repository.entity_deletes == [ + (collection_name(settings, scope, PROJECT), 77), + (collection_name(settings, scope, OTHER_PROJECT), 78), + ] @pytest.mark.asyncio @@ -446,18 +478,19 @@ async def test_reconciliation_deletes_only_absent_stable_keys( stale_key = VectorKey(entity_id=2, chunk_key="stale") await index.upsert( + PROJECT, [ VectorRecord(key=live_key, source_hash="a", values=(1.0, 0.0, 0.0)), VectorRecord(key=stale_key, source_hash="b", values=(0.0, 1.0, 0.0)), - ] + ], ) _, stored_records = repository.upserts[0] repository.ids = [record.record_id for record in stored_records] - await index.delete_orphans([live_key]) + await index.delete_orphans(PROJECT, [live_key]) assert repository.id_deletes == [ - (collection_name(settings, scope), [stored_records[1].record_id]) + (collection_name(settings, scope, PROJECT), [stored_records[1].record_id]) ] @@ -469,7 +502,7 @@ async def test_reconciliation_is_noop_without_orphans( repository = FakeRepository(dimensions=scope.dimensions) index = _index(scope, settings, repository) - await index.delete_orphans([]) + await index.delete_orphans(PROJECT, []) assert repository.id_deletes == [] @@ -483,11 +516,14 @@ async def test_reconciliation_deletes_orphans_incrementally( repository.ids = [f"orphan-{index}" for index in range(600)] index = _index(scope, settings, repository) - await index.delete_orphans([]) + await index.delete_orphans(PROJECT, []) assert [len(record_ids) for _, record_ids in repository.id_deletes] == [256, 256, 88] +# --- Search --- + + @pytest.mark.asyncio async def test_search_clamps_milvus_cosine_scores_and_orders_ties( scope: VectorIndexScope, @@ -502,25 +538,60 @@ async def test_search_clamps_milvus_cosine_scores_and_orders_ties( ] index = _index(scope, settings, repository) - matches = await index.search([1.0, 0.0, 0.0], limit=4) + matches = await index.search(QUERY, limit=4, projects=PROJECTS) assert [match.similarity for match in matches] == [1.0, 1.0, 0.1, 0.0] assert [match.key.entity_id for match in matches] == [0, 1, 2, 3] - assert repository.searches == [(collection_name(settings, scope), [1.0, 0.0, 0.0], 4)] + assert repository.searches == [(collection_name(settings, scope, PROJECT), QUERY, 4)] + + +@pytest.mark.asyncio +async def test_search_merges_the_collections_in_scope( + scope: VectorIndexScope, + settings: MilvusSettings, +) -> None: + """Milvus has no cross-collection search, so a wider scope merges per-project answers.""" + repository = FakeRepository(dimensions=scope.dimensions) + first = collection_name(settings, scope, PROJECT) + second = collection_name(settings, scope, OTHER_PROJECT) + repository.matches_by_collection = { + first: [ + MilvusStoredMatch(entity_id=1, chunk_key="a", score=0.9), + MilvusStoredMatch(entity_id=2, chunk_key="a", score=0.3), + ], + second: [ + MilvusStoredMatch(entity_id=3, chunk_key="a", score=0.8), + MilvusStoredMatch(entity_id=4, chunk_key="a", score=0.7), + ], + } + index = _index(scope, settings, repository) + + matches = await index.search(QUERY, limit=3, projects=ProjectScope.of([OTHER_PROJECT, PROJECT])) + + assert [(match.key.entity_id, match.similarity) for match in matches] == [ + (1, 0.9), + (3, 0.8), + (4, 0.7), + ] + # Each project's collection is asked for its own top ``limit`` before the merge. + assert repository.searches == [(first, QUERY, 3), (second, QUERY, 3)] + # Both collections were validated before being searched. + assert repository.loaded == [first, second] @pytest.mark.asyncio -async def test_empty_operations_do_not_initialize( +async def test_empty_operations_do_not_touch_milvus( scope: VectorIndexScope, settings: MilvusSettings, ) -> None: repository = FakeRepository() index = _index(scope, settings, repository) - await index.upsert([]) - await index.delete([]) - assert await index.search([], limit=10) == [] - assert await index.search([1.0, 0.0, 0.0], limit=0) == [] + await index.upsert(PROJECT, []) + await index.delete(PROJECT, []) + assert await index.search([], limit=10, projects=PROJECTS) == [] + assert await index.search(QUERY, limit=0, projects=PROJECTS) == [] + assert await index.search(QUERY, limit=10, projects=ProjectScope.of([])) == [] assert repository.closed == 0 @@ -535,7 +606,6 @@ def test_first_party_factory_loads_milvus() -> None: name, index = create_semantic_vector_index( session_maker=session_maker, - project_id=42, app_config=app_config, database_backend=DatabaseBackend.POSTGRES, embedding_provider=StubEmbeddingProvider(), @@ -545,4 +615,4 @@ def test_first_party_factory_loads_milvus() -> None: assert isinstance(index, MilvusVectorIndex) assert isinstance(index, SemanticVectorIndex) assert isinstance(index, SemanticVectorIndexReconciler) - assert index.scope.project_id == 42 + assert index.scope.dimensions == 3 diff --git a/tests/repository/test_pgvector_index.py b/tests/repository/test_pgvector_index.py index e1fa8bfe8..c46abca42 100644 --- a/tests/repository/test_pgvector_index.py +++ b/tests/repository/test_pgvector_index.py @@ -10,6 +10,7 @@ from basic_memory.repository import pgvector_index as pgvector_index_module from basic_memory.repository.pgvector_index import PgVectorIndex +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError from basic_memory.repository.semantic_vector_index import ( VectorDeletion, @@ -93,10 +94,13 @@ async def commit(self) -> None: self.commit_count += 1 +PROJECT = 7 +PROJECTS = ProjectScope.single(PROJECT) + + def _scope(dimensions: int = 4) -> VectorIndexScope: return VectorIndexScope( namespace="basic-memory-test", - project_id=7, embedding_identity="stub:4", dimensions=dimensions, ) @@ -208,10 +212,11 @@ async def test_upsert_resolves_stable_keys_and_writes_one_batch(monkeypatch) -> index._initialized = True await index.upsert( + PROJECT, [ VectorRecord(key=key_a, source_hash="hash-a", values=(1.0, 0.0, 0.0, 0.0)), VectorRecord(key=key_b, source_hash="hash-b", values=(0.0, 1.0, 0.0, 0.0)), - ] + ], ) insert_call = next(call for call in session.calls if "INSERT INTO" in call[0]) @@ -246,7 +251,9 @@ async def test_upsert_skips_stale_source_generation(monkeypatch) -> None: index = PgVectorIndex(MagicMock(), _scope()) index._initialized = True - await index.upsert([VectorRecord(key=key, source_hash="old-hash", values=(1.0, 0.0, 0.0, 0.0))]) + await index.upsert( + PROJECT, [VectorRecord(key=key, source_hash="old-hash", values=(1.0, 0.0, 0.0, 0.0))] + ) lock_call = next(call for call in session.calls if "SELECT id, entity_id" in call[0]) assert "FOR UPDATE" in lock_call[0] @@ -263,7 +270,9 @@ async def test_upsert_rejects_missing_manifest_key(monkeypatch) -> None: index._initialized = True with pytest.raises(RuntimeError, match="manifest rows are missing"): - await index.upsert([VectorRecord(key=key, source_hash="hash", values=(1.0, 0.0, 0.0, 0.0))]) + await index.upsert( + PROJECT, [VectorRecord(key=key, source_hash="hash", values=(1.0, 0.0, 0.0, 0.0))] + ) @pytest.mark.asyncio @@ -274,10 +283,10 @@ async def test_delete_stable_keys_and_entity(monkeypatch) -> None: index = PgVectorIndex(MagicMock(), _scope()) index._initialized = True - await index.delete([]) - await index.delete([VectorDeletion(key=key, source_hash="hash")]) - await index.delete_entity(13) - await index.delete_orphans([key]) + await index.delete(PROJECT, []) + await index.delete(PROJECT, [VectorDeletion(key=key, source_hash="hash")]) + await index.delete_entity(PROJECT, 13) + await index.delete_orphans(PROJECT, [key]) delete_lock = next(call for call in session.calls if "SELECT id, entity_id" in call[0]) assert "source_hash = :source_hash_0" in delete_lock[0] @@ -307,12 +316,14 @@ async def test_search_returns_normalized_stable_matches(monkeypatch) -> None: index = PgVectorIndex(MagicMock(), _scope()) index._initialized = True - assert await index.search([], limit=5) == [] - assert await index.search([1.0, 0.0, 0.0, 0.0], limit=0) == [] + assert await index.search([], limit=5, projects=PROJECTS) == [] + assert await index.search([1.0, 0.0, 0.0, 0.0], limit=0, projects=PROJECTS) == [] + assert await index.search([1.0, 0.0, 0.0, 0.0], limit=5, projects=ProjectScope.of([])) == [] with pytest.raises(ValueError, match="expected 4, got 2"): - await index.search([1.0, 0.0], limit=5) + await index.search([1.0, 0.0], limit=5, projects=PROJECTS) + assert session.calls == [] - matches = await index.search([1.0, 0.0, 0.0, 0.0], limit=5) + matches = await index.search([1.0, 0.0, 0.0, 0.0], limit=5, projects=PROJECTS) assert [(match.key.entity_id, match.similarity) for match in matches] == [ (14, 1.0), @@ -322,8 +333,26 @@ async def test_search_returns_normalized_stable_matches(monkeypatch) -> None: assert "c.entity_id ASC, c.chunk_key ASC" in search_call[0] assert search_call[1] == { "query": "[1,0,0,0]", - "project_id": 7, + "scope_0": 7, "dimensions": 4, "embedding_identity": "stub:4", "limit": 5, } + + +@pytest.mark.asyncio +async def test_search_binds_every_project_in_scope(monkeypatch) -> None: + """A multi-project scope filters both the embedding and manifest rows by the same ids.""" + session = FakeSession(search_rows=[]) + _install_session(monkeypatch, session) + index = PgVectorIndex(MagicMock(), _scope()) + index._initialized = True + + await index.search([1.0, 0.0, 0.0, 0.0], limit=5, projects=ProjectScope.of([9, 7])) + + sql, params = next(call for call in session.calls if "AS similarity" in call[0]) + assert params is not None + assert "WHERE e.project_id IN (:scope_0, :scope_1)" in sql + assert "AND c.project_id IN (:scope_0, :scope_1)" in sql + assert params["scope_0"] == 7 + assert params["scope_1"] == 9 diff --git a/tests/repository/test_search_reader.py b/tests/repository/test_search_reader.py index ed6da6e62..cccf2668b 100644 --- a/tests/repository/test_search_reader.py +++ b/tests/repository/test_search_reader.py @@ -233,7 +233,7 @@ async def test_built_in_adapter_reads_manifest_readiness_when_tracing(monkeypatc assert trace.readiness is readiness read_readiness.assert_awaited_once_with(session, SCOPE, "sqlite-vec", "fake:384") - adapter.search.assert_awaited_once_with([0.1], limit=5) + adapter.search.assert_awaited_once_with([0.1], limit=5, projects=SCOPE) @pytest.mark.asyncio diff --git a/tests/repository/test_search_trace.py b/tests/repository/test_search_trace.py index 89efb107d..055050960 100644 --- a/tests/repository/test_search_trace.py +++ b/tests/repository/test_search_trace.py @@ -68,11 +68,10 @@ def runtime_log_attrs(self) -> dict[str, Any]: class _TraceVectorIndex: - def __init__(self, project_id: int) -> None: + def __init__(self) -> None: self.matches: list[VectorMatch] = [] self.scope = VectorIndexScope( namespace="trace-test", - project_id=project_id, embedding_identity="trace-embedding", dimensions=4, ) @@ -80,16 +79,18 @@ def __init__(self, project_id: int) -> None: async def initialize(self) -> None: return None - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: return None - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: return None - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: return None - async def search(self, query: Sequence[float], *, limit: int) -> list[VectorMatch]: + async def search( + self, query: Sequence[float], *, limit: int, projects: ProjectScope + ) -> list[VectorMatch]: return self.matches[:limit] @@ -525,7 +526,7 @@ def _repository( "semantic_vector_k": 10, } ) - vector_index = _TraceVectorIndex(test_project.id) + vector_index = _TraceVectorIndex() repository_type = ( PostgresSearchRepository if config.database_backend == DatabaseBackend.POSTGRES diff --git a/tests/repository/test_semantic_search_base.py b/tests/repository/test_semantic_search_base.py index 4a5857aca..43fe5806f 100644 --- a/tests/repository/test_semantic_search_base.py +++ b/tests/repository/test_semantic_search_base.py @@ -148,7 +148,6 @@ class _RecordingVectorIndex: scope = VectorIndexScope( namespace="basic-memory-test", - project_id=1, embedding_identity="stub:4", dimensions=4, ) @@ -159,13 +158,13 @@ def __init__(self) -> None: async def initialize(self) -> None: return None - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: self.upserted_records.extend(records) - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: return None - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: return None async def search( @@ -173,6 +172,7 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: return [] @@ -475,7 +475,9 @@ async def test_external_reconciliation_holds_project_lock_through_orphan_cleanup events: list[str] = [] adapter: Any = SimpleNamespace( scope=_RecordingVectorIndex.scope, - delete_orphans=AsyncMock(side_effect=lambda _live_keys: events.append("delete_orphans")), + delete_orphans=AsyncMock( + side_effect=lambda _project_id, _live_keys: events.append("delete_orphans") + ), ) repo._semantic_vector_index = adapter session = AsyncMock() @@ -508,7 +510,7 @@ async def fake_scoped_session(_session_maker): assert events == ["project_lock", "manifest_read", "delete_orphans", "commit"] adapter.delete_orphans.assert_awaited_once_with( - [VectorKey(entity_id=41, chunk_key="entity:41:0")] + 1, [VectorKey(entity_id=41, chunk_key="entity:41:0")] ) @@ -654,7 +656,9 @@ async def test_project_vector_cleanup_uses_available_adapter( events: list[str] = [] adapter: Any = SimpleNamespace( initialize=AsyncMock(side_effect=lambda: events.append("initialize")), - delete_entity=AsyncMock(side_effect=lambda _entity_id: events.append("delete")), + delete_entity=AsyncMock( + side_effect=lambda _project_id, _entity_id: events.append("delete") + ), ) repo._semantic_vector_index = adapter repo._semantic_vector_index_name = "milvus" @@ -697,8 +701,8 @@ async def fake_scoped_session(_session_maker): adapter.initialize.assert_awaited_once() assert adapter.delete_entity.await_args_list == [ - ((41,), {}), - ((42,), {}), + ((1, 41), {}), + ((1, 42), {}), ] expected_events = [ "project_lock", @@ -912,7 +916,9 @@ async def test_external_entity_cleanup_uses_matching_project_adapter(monkeypatch events: list[str] = [] adapter: Any = SimpleNamespace( initialize=AsyncMock(side_effect=lambda: events.append("initialize")), - delete_entity=AsyncMock(side_effect=lambda _entity_id: events.append("delete")), + delete_entity=AsyncMock( + side_effect=lambda _project_id, _entity_id: events.append("delete") + ), ) repo._semantic_vector_index = adapter repo._semantic_vector_index_name = "milvus" @@ -948,7 +954,7 @@ async def fake_scoped_session(_session_maker): ) adapter.initialize.assert_awaited_once() - assert adapter.delete_entity.await_args_list == [((41,), {}), ((42,), {})] + assert adapter.delete_entity.await_args_list == [((1, 41), {}), ((1, 42), {})] assert events == [ "project_lock", "ownership_read", diff --git a/tests/repository/test_semantic_vector_index.py b/tests/repository/test_semantic_vector_index.py index 96ec7b4f8..224ab0dd9 100644 --- a/tests/repository/test_semantic_vector_index.py +++ b/tests/repository/test_semantic_vector_index.py @@ -14,6 +14,7 @@ from basic_memory.repository.embedding_provider import EmbeddingProvider from basic_memory.repository.postgres_search_repository import PostgresSearchRepository from basic_memory.repository.search_repository import create_search_repository +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_errors import ( SemanticDependenciesMissingError, ) @@ -59,13 +60,13 @@ def __init__(self, scope: VectorIndexScope): async def initialize(self) -> None: return None - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: return None - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: return None - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: return None async def search( @@ -73,6 +74,7 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: return [] @@ -92,7 +94,6 @@ def _postgres_config(**overrides: object) -> BasicMemoryConfig: def test_vector_contract_values_and_dimension_validation() -> None: scope = VectorIndexScope( namespace="basic-memory-test", - project_id=7, embedding_identity="stub:3", dimensions=3, ) @@ -127,9 +128,9 @@ def test_selector_defaults_to_pgvector_and_sqlite_remains_automatic() -> None: _postgres_config(semantic_vector_index="test-extension") -def test_scope_is_stable_credential_free_and_project_isolated() -> None: +def test_scope_is_stable_and_credential_free() -> None: provider: EmbeddingProvider = StubEmbeddingProvider() - first = build_vector_index_scope(_postgres_config(), provider, project_id=7) + first = build_vector_index_scope(_postgres_config(), provider) rotated_password = build_vector_index_scope( _postgres_config( database_url=( @@ -137,15 +138,12 @@ def test_scope_is_stable_credential_free_and_project_isolated() -> None: ) ), provider, - project_id=7, ) - other_project = build_vector_index_scope(_postgres_config(), provider, project_id=8) other_user = build_vector_index_scope( _postgres_config( database_url="postgresql+asyncpg://tenant-user:secret@db.example.test:5432/memory" ), provider, - project_id=7, ) other_schema = build_vector_index_scope( _postgres_config( @@ -155,7 +153,6 @@ def test_scope_is_stable_credential_free_and_project_isolated() -> None: ) ), provider, - project_id=7, ) first_socket = build_vector_index_scope( _postgres_config( @@ -164,7 +161,6 @@ def test_scope_is_stable_credential_free_and_project_isolated() -> None: ) ), provider, - project_id=7, ) other_socket = build_vector_index_scope( _postgres_config( @@ -173,18 +169,16 @@ def test_scope_is_stable_credential_free_and_project_isolated() -> None: ) ), provider, - project_id=7, ) assert first.namespace == rotated_password.namespace assert "secret" not in first.namespace - assert first.project_id != other_project.project_id assert first.namespace != other_user.namespace assert first.namespace != other_schema.namespace assert first_socket.namespace != other_socket.namespace assert first.embedding_identity == "StubEmbeddingProvider:stub-model:3" assert first.dimensions == 3 - assert first.storage_key == rotated_password.storage_key + assert first == rotated_password def test_milvus_without_optional_dependencies_reports_install_extra(monkeypatch) -> None: @@ -213,7 +207,6 @@ def import_without_pymilvus( ): create_semantic_vector_index( session_maker=MagicMock(), - project_id=7, app_config=config, database_backend=DatabaseBackend.POSTGRES, embedding_provider=StubEmbeddingProvider(), @@ -222,7 +215,7 @@ def import_without_pymilvus( def test_search_repository_composition_root_injects_selected_adapter(monkeypatch) -> None: provider = StubEmbeddingProvider() - scope = build_vector_index_scope(_postgres_config(), provider, project_id=7) + scope = build_vector_index_scope(_postgres_config(), provider) index = StubVectorIndex(scope) monkeypatch.setattr( "basic_memory.repository.search_repository.create_embedding_provider", diff --git a/tests/repository/test_sqlite_vector_search_repository.py b/tests/repository/test_sqlite_vector_search_repository.py index f981fac1c..c53291b26 100644 --- a/tests/repository/test_sqlite_vector_search_repository.py +++ b/tests/repository/test_sqlite_vector_search_repository.py @@ -18,6 +18,7 @@ from basic_memory.repository.prefixing_provider import PrefixingEmbeddingProvider from basic_memory.repository import search_repository_base as search_repository_base_module from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_errors import SemanticVectorIndexExtensionError from basic_memory.repository.semantic_vector_index import ( VectorDeletion, @@ -76,7 +77,6 @@ class RecordingVectorIndex: def __init__(self) -> None: self.scope = VectorIndexScope( namespace="test", - project_id=1, embedding_identity="test", dimensions=4, ) @@ -91,20 +91,20 @@ def __init__(self) -> None: async def initialize(self) -> None: return None - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: self.upsert_calls.append(list(records)) if self.fail_upsert: raise RuntimeError("adapter write failed") self.records.update({record.key: record.values for record in records}) - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: self.deleted_entities.extend(sorted({record.key.entity_id for record in records})) if self.fail_delete_entity: raise RuntimeError("adapter delete failed") for record in records: self.records.pop(record.key, None) - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: self.deleted_entities.append(entity_id) if self.fail_delete_entity: raise RuntimeError("adapter delete failed") @@ -112,7 +112,7 @@ async def delete_entity(self, entity_id: int) -> None: key: values for key, values in self.records.items() if key.entity_id != entity_id } - async def delete_orphans(self, live_keys: Sequence[VectorKey]) -> None: + async def delete_orphans(self, project_id: int, live_keys: Sequence[VectorKey]) -> None: self.reconcile_calls.append(list(live_keys)) live_key_set = set(live_keys) self.records = {key: values for key, values in self.records.items() if key in live_key_set} @@ -122,6 +122,7 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: if self.fail_search: raise RuntimeError("adapter query failed") @@ -436,7 +437,7 @@ async def test_sqlite_vec_reconciliation_is_project_scoped(search_repository): ) await session.commit() - await index.delete_orphans([]) + await index.delete_orphans(search_repository.project_id, []) async with db.scoped_session(search_repository.session_maker) as session: remaining = await session.execute( @@ -448,6 +449,71 @@ async def test_sqlite_vec_reconciliation_is_project_scoped(search_repository): assert remaining.scalars().all() == [902, 903] +@pytest.mark.asyncio +async def test_sqlite_vec_search_reads_every_project_in_scope(search_repository): + """One statement answers a multi-project scope; a single-project scope stays isolated.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec search behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await search_repository.init_search_index() + index = cast(SQLiteVecIndex, search_repository._semantic_vector_index) + embedding_identity = search_repository._embedding_model_key() + own_project = search_repository.project_id + other_project = own_project + 1 + + async with db.scoped_session(search_repository.session_maker) as session: + await index._ensure_loaded(session) + await session.execute( + text( + "INSERT INTO search_vector_chunks (" + "id, entity_id, project_id, chunk_key, chunk_text, source_hash, " + "entity_fingerprint, embedding_model, vector_index, embedding_status" + ") VALUES (" + ":id, :entity_id, :project_id, :chunk_key, 'text', 'hash', " + "'fingerprint', :embedding_model, 'sqlite-vec', 'ready')" + ), + [ + { + "id": 911, + "entity_id": 911, + "project_id": own_project, + "chunk_key": "entity:911:0", + "embedding_model": embedding_identity, + }, + { + "id": 912, + "entity_id": 912, + "project_id": other_project, + "chunk_key": "entity:912:0", + "embedding_model": embedding_identity, + }, + ], + ) + await session.execute( + text( + "INSERT INTO search_vector_embeddings (rowid, embedding, source_hash) " + "VALUES (:rowid, :embedding, 'hash')" + ), + [ + {"rowid": 911, "embedding": "[1,0,0,0]"}, + {"rowid": 912, "embedding": "[0,1,0,0]"}, + ], + ) + await session.commit() + + query = [1.0, 0.0, 0.0, 0.0] + both = await index.search( + query, limit=10, projects=ProjectScope.of([other_project, own_project]) + ) + own_only = await index.search(query, limit=10, projects=ProjectScope.single(own_project)) + nothing = await index.search(query, limit=10, projects=ProjectScope.of([])) + + assert [match.key.entity_id for match in both] == [911, 912] + assert [match.key.entity_id for match in own_only] == [911] + assert nothing == [] + + @pytest.mark.asyncio async def test_sqlite_vec_delete_requires_pending_source_generation(search_repository): """A stale delete cannot remove a same-source vector that is already ready.""" @@ -486,7 +552,7 @@ async def test_sqlite_vec_delete_requires_pending_source_generation(search_repos await session.commit() deletion = VectorDeletion(key=key, source_hash="hash") - await index.delete([deletion]) + await index.delete(search_repository.project_id, [deletion]) async with db.scoped_session(search_repository.session_maker) as session: assert ( await session.scalar( @@ -499,7 +565,7 @@ async def test_sqlite_vec_delete_requires_pending_source_generation(search_repos ) await session.commit() - await index.delete([deletion]) + await index.delete(search_repository.project_id, [deletion]) async with db.scoped_session(search_repository.session_maker) as session: vector_count = await session.scalar( text("SELECT COUNT(*) FROM search_vector_embeddings WHERE rowid = 907") @@ -1459,19 +1525,19 @@ async def fake_scoped_session(_session_maker): monkeypatch.setattr(index, "_ensure_loaded", AsyncMock()) query_embedding = [0.1] * search_repository._vector_dimensions - await index.search(query_embedding, limit=10000) + await index.search(query_embedding, limit=10000, projects=search_repository.scope) assert captured_params == [ { "query": "[0.1, 0.1, 0.1, 0.1]", "vector_k": SQLITE_VEC_MAX_K, - "project_id": search_repository.project_id, + "scope_0": search_repository.project_id, "embedding_identity": search_repository._embedding_model_key(), "limit": 10000, } ] captured_params.clear() - await index.search(query_embedding, limit=500) + await index.search(query_embedding, limit=500, projects=search_repository.scope) assert captured_params[0]["vector_k"] == 500 assert captured_params[0]["limit"] == 500 diff --git a/tests/repository/test_vector_manifest_generation_ownership.py b/tests/repository/test_vector_manifest_generation_ownership.py index f92abbd0d..62f502476 100644 --- a/tests/repository/test_vector_manifest_generation_ownership.py +++ b/tests/repository/test_vector_manifest_generation_ownership.py @@ -13,6 +13,7 @@ from basic_memory.config import BasicMemoryConfig, DatabaseBackend from basic_memory.repository.postgres_search_repository import PostgresSearchRepository from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.semantic_vector_index import ( VectorDeletion, VectorIndexScope, @@ -74,17 +75,17 @@ def scope(self) -> VectorIndexScope: async def initialize(self) -> None: return None - async def upsert(self, records: Sequence[VectorRecord]) -> None: + async def upsert(self, project_id: int, records: Sequence[VectorRecord]) -> None: for record in records: self.records[record.key] = record - async def delete(self, records: Sequence[VectorDeletion]) -> None: + async def delete(self, project_id: int, records: Sequence[VectorDeletion]) -> None: for deletion in records: current = self.records.get(deletion.key) if current is not None and current.source_hash == deletion.source_hash: self.records.pop(deletion.key) - async def delete_entity(self, entity_id: int) -> None: + async def delete_entity(self, project_id: int, entity_id: int) -> None: self.records = { key: record for key, record in self.records.items() if key.entity_id != entity_id } @@ -94,6 +95,7 @@ async def search( query: Sequence[float], *, limit: int, + projects: ProjectScope, ) -> list[VectorMatch]: return [] @@ -144,7 +146,6 @@ async def _repositories( vector_index = InMemoryExternalVectorIndex( VectorIndexScope( namespace="generation-ownership", - project_id=project_id, embedding_identity="test:4", dimensions=4, )