From 2c7097ab05ab1abffdbed94c3189ad27463bf453 Mon Sep 17 00:00:00 2001 From: phernandez Date: Tue, 15 Sep 2026 14:29:05 -0500 Subject: [PATCH] refactor(core): bind vector adapters to the database and search them over a ProjectScope The vector adapters were bound to one project through VectorIndexScope.project_id, which made a multi-project search impossible to express: the reader could only ask the adapter for the project it was built for. The project was never part of the vector's identity (entity ids are database-wide primary keys); it is the partition an operation touches. VectorIndexScope is now the database namespace plus the embedding schema. Every write names its project (upsert, delete, delete_entity, delete_orphans take project_id first) and search takes a ProjectScope, so one sqlite-vec or pgvector adapter answers any set of projects with one statement through the scope's IN predicate. An empty scope returns nothing without touching storage. Milvus keeps a collection per project, validates each collection on its first use instead of in initialize(), and searches the collections in scope, merging by similarity. SemanticSearch passes its own scope to the adapter, so a project repository's vector search is unchanged. The repositories build the scope without a project and the factory no longer takes one. Two dead lookup helpers in the built-in adapters go with the change. Test doubles implement the new signatures; the pgvector, sqlite-vec, and Milvus suites gain multi-project and empty-scope cases. No query behavior changes. Part of #1558. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_019YW9ysxugGGBCNEGzsxtFV Signed-off-by: phernandez --- CHANGELOG.md | 9 + src/basic_memory/repository/milvus_index.py | 128 ++++++------ src/basic_memory/repository/pgvector_index.py | 75 +++---- .../repository/postgres_search_repository.py | 12 +- src/basic_memory/repository/search_reader.py | 8 +- .../repository/search_repository.py | 1 - .../repository/search_repository_base.py | 14 +- .../repository/semantic_vector_index.py | 40 ++-- .../semantic_vector_index_factory.py | 10 +- .../repository/sqlite_search_repository.py | 12 +- .../repository/sqlite_vec_index.py | 68 ++---- test-int/semantic/conftest.py | 1 - test-int/semantic/test_milvus_lite.py | 33 +-- test-int/test_embedding_status_vec0.py | 6 +- tests/repository/test_milvus_index.py | 194 ++++++++++++------ tests/repository/test_pgvector_index.py | 55 +++-- tests/repository/test_search_reader.py | 2 +- tests/repository/test_search_trace.py | 15 +- tests/repository/test_semantic_search_base.py | 28 ++- .../repository/test_semantic_vector_index.py | 25 +-- .../test_sqlite_vector_search_repository.py | 88 +++++++- ...st_vector_manifest_generation_ownership.py | 9 +- 22 files changed, 484 insertions(+), 349 deletions(-) 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, )