From cb85685af137967013710b6f69fabd60e78116c8 Mon Sep 17 00:00:00 2001 From: phernandez Date: Tue, 15 Sep 2026 16:03:17 -0500 Subject: [PATCH] feat(api): search an explicit set of projects with one query QUERY /v2/search/ (and POST /v2/search/ for clients that cannot send QUERY) searches an explicit set of projects in one database. The body is the project search body plus project_ids, a required list of internal ids the caller has already authorized; an empty list answers no rows and there is no spelling for every project. create_search_reader composes a SearchReader over a ProjectScope with no project repository, resolving the same shared embedding provider, vector adapter, and reranker a repository gets, so one project through this route ranks exactly as its own route does and several projects are one ranking over the union. ScopedSearchService runs the query and hydrates only from projects in scope: owning entities, relation endpoints, project identities, and valid-time assertions are all read under the scope predicate. The query preparation the project service already had (criteria normalization, legacy note-type expansion, relaxed full-text eligibility) moves to module functions both services call, with the note-type expansion now bounded to the scope it serves. Every SearchResult, on both routes, carries project_id and project_external_id. The two search routes share one error boundary. An empty scope returns before embedding. The sqlite-vec adapter converts a missing enable_load_extension into the same typed dependency error the repository raises, so a scoped vector search on such a host answers 400 rather than 500. Part of #1558. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_019YW9ysxugGGBCNEGzsxtFV Signed-off-by: phernandez --- CHANGELOG.md | 10 + src/basic_memory/api/app.py | 3 + src/basic_memory/api/v2/routers/__init__.py | 2 + .../api/v2/routers/scoped_search_router.py | 114 +++ .../api/v2/routers/search_router.py | 37 +- src/basic_memory/api/v2/utils.py | 44 +- src/basic_memory/repository/search_reader.py | 5 + .../repository/search_repository.py | 63 ++ .../repository/sqlite_vec_index.py | 13 +- src/basic_memory/schemas/search.py | 19 +- .../services/scoped_search_service.py | 156 +++++ src/basic_memory/services/search_service.py | 241 ++++--- tests/api/v2/test_scoped_search_router.py | 657 ++++++++++++++++++ tests/api/v2/test_search_router_telemetry.py | 5 +- tests/repository/test_search_reader.py | 17 + .../repository/test_search_reader_factory.py | 115 +++ tests/repository/test_sqlite_vec_index.py | 41 ++ 17 files changed, 1408 insertions(+), 134 deletions(-) create mode 100644 src/basic_memory/api/v2/routers/scoped_search_router.py create mode 100644 src/basic_memory/services/scoped_search_service.py create mode 100644 tests/api/v2/test_scoped_search_router.py create mode 100644 tests/repository/test_search_reader_factory.py create mode 100644 tests/repository/test_sqlite_vec_index.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 0d347cd53..a53634749 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,16 @@ ### Features +- **#1558**: `QUERY /v2/search/` (and `POST /v2/search/` for clients that cannot send + QUERY) searches an explicit set of projects in one database with one query. The body + is the project search body plus `project_ids`, a required list of internal ids the + caller has already authorized; an empty list answers no rows and there is no way to + ask for every project. Full-text, vector, and hybrid retrieval run the same reader + the project route runs, so one project here ranks exactly as its own route does, and + results are one ranking over the union rather than merged per-project pages. Hits are + hydrated only from projects in scope. Every search result, on both routes, now + carries `project_id` and `project_external_id`. + - **#1512**: Word, PowerPoint, and CSV files get the same sidecar Markdown note a PDF gets. `bm import document ` indexes the project, extracts the file, and writes `..md` next to it plus a run note under diff --git a/src/basic_memory/api/app.py b/src/basic_memory/api/app.py index 5e2cb8cf6..c2779448b 100644 --- a/src/basic_memory/api/app.py +++ b/src/basic_memory/api/app.py @@ -14,6 +14,7 @@ knowledge_router as v2_knowledge, project_router as v2_project, memory_router as v2_memory, + scoped_search_router as v2_scoped_search, search_router as v2_search, resource_router as v2_resource, directory_router as v2_directory, @@ -136,6 +137,8 @@ async def workspace_permalink_context_middleware(request: Request, call_next): app.include_router(v2_schema, prefix="/v2/projects/{project_id}") app.include_router(v2_inspect, prefix="/v2/projects/{project_id}") app.include_router(v2_project, prefix="/v2") +# Database-scoped search: one query over an explicit set of projects. +app.include_router(v2_scoped_search, prefix="/v2") # Legacy web app proxy paths (compat with /proxy/projects/projects) app.include_router(v2_project, prefix="/proxy/projects") diff --git a/src/basic_memory/api/v2/routers/__init__.py b/src/basic_memory/api/v2/routers/__init__.py index 29015ba37..9c0bc36ab 100644 --- a/src/basic_memory/api/v2/routers/__init__.py +++ b/src/basic_memory/api/v2/routers/__init__.py @@ -6,6 +6,7 @@ from basic_memory.api.v2.routers.knowledge_router import router as knowledge_router from basic_memory.api.v2.routers.project_router import router as project_router from basic_memory.api.v2.routers.memory_router import router as memory_router +from basic_memory.api.v2.routers.scoped_search_router import router as scoped_search_router from basic_memory.api.v2.routers.search_router import router as search_router from basic_memory.api.v2.routers.resource_router import router as resource_router from basic_memory.api.v2.routers.directory_router import router as directory_router @@ -22,6 +23,7 @@ "knowledge_router", "project_router", "memory_router", + "scoped_search_router", "search_router", "resource_router", "directory_router", diff --git a/src/basic_memory/api/v2/routers/scoped_search_router.py b/src/basic_memory/api/v2/routers/scoped_search_router.py new file mode 100644 index 000000000..a6fe02a54 --- /dev/null +++ b/src/basic_memory/api/v2/routers/scoped_search_router.py @@ -0,0 +1,114 @@ +"""Database-scoped search: one query over an explicit set of projects. + +The caller names the projects. Nothing here discovers a scope, and an empty set +answers no rows; Cloud resolves authorization first and passes the effective ids. +The route shares the project route's reader, so one project here ranks exactly as +that project's own route does, and several projects are one ranking over the union +rather than per-project pages merged afterwards. +""" + +import asyncio + +import logfire +from fastapi import APIRouter, Query, Response + +from basic_memory import db +from basic_memory.api.v2.utils import ( + load_temporal_metadata, + search_error_boundary, + to_search_results, +) +from basic_memory.deps import AppConfigDep, SessionMakerDep +from basic_memory.repository.search_repository import create_search_reader +from basic_memory.repository.search_scope import ProjectScope +from basic_memory.schemas.search import ScopedSearchQuery, SearchResponse, SearchRetrievalMode +from basic_memory.services.scoped_search_service import ScopedSearchService + +# App registration mounts this router at /v2. +router = APIRouter(tags=["search"]) + + +@router.api_route( + "/search/", + methods=["QUERY"], + response_model=SearchResponse, + include_in_schema=False, +) +@router.post("/search/", response_model=SearchResponse) +async def search_scope( + query: ScopedSearchQuery, + app_config: AppConfigDep, + session_maker: SessionMakerDep, + response: Response, + page: int = Query(1, ge=1), + page_size: int = Query(10, ge=1, le=1000), +) -> SearchResponse: + """Search an explicit set of projects in this database. + + Results are one ranking across the whole scope and carry ``project_id`` and + ``project_external_id``. Hydration reads only from projects in scope. + """ + response.headers["Accept-Query"] = "application/json" + # The read cache is keyed by one project's generation. A set of projects has no + # single generation to invalidate on, so this route is never cached. + response.headers["Cache-Control"] = "no-store" + + scope = ProjectScope.of(query.project_ids) + service = ScopedSearchService( + session_maker, scope, create_search_reader(session_maker, scope, app_config) + ) + temporal_requested = query.has_temporal_filter() + exact_count_available = query.retrieval_mode == SearchRetrievalMode.FTS + offset = (page - 1) * page_size + + with logfire.span( + "api.request.search_scope", + entrypoint="api", + domain="search", + action="search_scope", + project_count=len(scope.project_ids), + page=page, + page_size=page_size, + retrieval_mode=query.retrieval_mode.value, + has_temporal_filter=temporal_requested, + ): + with search_error_boundary(): + if exact_count_available: + results, total = await asyncio.gather( + service.search(query, limit=page_size, offset=offset), + service.count(query), + ) + has_more = offset + len(results) < total + else: + # Trigger: semantic modes would need another vector or hybrid pass to count. + # Why: a search should not pay for a second semantic retrieval. + # Outcome: probe one row past the page, leave total at 0, mark it inexact. + results = await service.search(query, limit=page_size + 1, offset=offset) + total = 0 + has_more = len(results) > page_size + results = results[:page_size] + + temporal_by_source = {} + project_external_ids: dict[int, str] = {} + if results: + async with db.scoped_session(session_maker) as session: + project_external_ids = await service.project_external_ids(session, results) + if temporal_requested: + temporal_by_source = await load_temporal_metadata(service, session, results) + search_results = await to_search_results( + service, + results, + temporal_by_source=temporal_by_source, + project_external_ids=project_external_ids, + ) + return SearchResponse( + results=search_results, + current_page=page, + page_size=page_size, + total=total, + total_is_exact=exact_count_available, + has_more=has_more, + # None, not False, when nothing was asked: an ordinary search payload stays + # exactly what it was before valid time existed. + temporal_applied=True if temporal_requested else None, + ) diff --git a/src/basic_memory/api/v2/routers/search_router.py b/src/basic_memory/api/v2/routers/search_router.py index 2e8460142..f5c137486 100644 --- a/src/basic_memory/api/v2/routers/search_router.py +++ b/src/basic_memory/api/v2/routers/search_router.py @@ -9,11 +9,15 @@ from contextlib import nullcontext from typing import Annotated -from fastapi import APIRouter, Depends, HTTPException, Path, Response +from fastapi import APIRouter, Depends, Path, Response import logfire from basic_memory import db -from basic_memory.api.v2.utils import load_temporal_metadata, to_search_results +from basic_memory.api.v2.utils import ( + load_temporal_metadata, + search_error_boundary, + to_search_results, +) from basic_memory.deps import ( EntityServiceV2ExternalDep, MemoryTimeIndexRepositoryV2ExternalDep, @@ -32,12 +36,6 @@ read_cache_request_digest, ) from basic_memory.read_cache.policy import SEARCH_READ_CACHE_TTL_SECONDS -from basic_memory.repository.semantic_errors import ( - RerankProviderContractError, - RerankTransientError, - SemanticDependenciesMissingError, - SemanticSearchDisabledError, -) from basic_memory.schemas.search import SearchQuery, SearchResponse, SearchRetrievalMode from basic_memory.services.search_guidance import unspaced_script_query_hint @@ -91,6 +89,7 @@ async def search( session_maker: SessionMakerDep, read_cache: SearchReadCacheDep, response: Response, + internal_project_id: ProjectExternalIdPathDep, project_id: str = Path(..., description="Project external UUID"), page: int = 1, page_size: int = 10, @@ -157,7 +156,7 @@ async def search( offset = (page - 1) * page_size exact_count_available = query.retrieval_mode == SearchRetrievalMode.FTS - try: + with search_error_boundary(): with logfire.span( "api.search.search.execute_query", domain="search", @@ -176,21 +175,6 @@ async def search( query, limit=page_size + 1, offset=offset ) total = 0 - except SemanticSearchDisabledError as exc: - raise HTTPException(status_code=400, detail=str(exc)) from exc - except SemanticDependenciesMissingError as exc: - raise HTTPException(status_code=400, detail=str(exc)) from exc - except RerankTransientError as exc: - # Returning raw retrieval order would make pagination inconsistent with - # earlier reranked pages. Preserve ordering semantics and make the outage - # explicitly retryable instead. - raise HTTPException(status_code=503, detail=str(exc)) from exc - except RerankProviderContractError as exc: - # Upstream reranker returned a malformed response — an upstream fault, not a - # client error and not a transient outage (those map to a retryable 503). - raise HTTPException(status_code=502, detail=str(exc)) from exc - except ValueError as exc: - raise HTTPException(status_code=400, detail=str(exc)) from exc with logfire.span( "api.search.search.paginate_results", @@ -228,7 +212,10 @@ async def search( temporal_repository, session, results ) search_results = await to_search_results( - entity_service, results, temporal_by_source=temporal_by_source + entity_service, + results, + temporal_by_source=temporal_by_source, + project_external_ids={internal_project_id: project_id}, ) with logfire.span( "api.search.search.build_response", diff --git a/src/basic_memory/api/v2/utils.py b/src/basic_memory/api/v2/utils.py index c9d0350bd..82c7942f9 100644 --- a/src/basic_memory/api/v2/utils.py +++ b/src/basic_memory/api/v2/utils.py @@ -1,11 +1,19 @@ from collections import defaultdict -from collections.abc import Mapping +from collections.abc import Iterator, Mapping +from contextlib import contextmanager from typing import Any, Protocol, Optional, List, Sequence import logfire +from fastapi import HTTPException from sqlalchemy.ext.asyncio import AsyncSession from basic_memory.models import MemoryTimeIndex from basic_memory.repository.search_repository import SearchIndexRow +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, + SemanticDependenciesMissingError, + SemanticSearchDisabledError, +) from basic_memory.schemas.memory import ( EntitySummary, ObservationSummary, @@ -54,6 +62,28 @@ async def find_for_sources( type TemporalMetadataBySource = Mapping[tuple[str, int], list[TemporalResultMetadata]] +@contextmanager +def search_error_boundary() -> Iterator[None]: + """Map search failures onto HTTP statuses the same way on every search route.""" + try: + yield + except SemanticSearchDisabledError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + except SemanticDependenciesMissingError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + except RerankTransientError as exc: + # Returning raw retrieval order would make pagination inconsistent with + # earlier reranked pages. Preserve ordering semantics and make the outage + # explicitly retryable instead. + raise HTTPException(status_code=503, detail=str(exc)) from exc + except RerankProviderContractError as exc: + # Upstream reranker returned a malformed response: an upstream fault, not a + # client error and not a transient outage (those map to a retryable 503). + raise HTTPException(status_code=502, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + async def get_entities_by_id_lookup( entity_service: EntityServiceBatchLookup, entity_ids: Sequence[int], @@ -304,7 +334,13 @@ async def to_search_results( results: List[SearchIndexRow], *, temporal_by_source: TemporalMetadataBySource | None = None, + project_external_ids: Mapping[int, str] | None = None, ) -> list[SearchResult]: + """Shape one page of search rows into the public result. + + ``project_external_ids`` maps each row's project to its external id; a caller + that knows the projects passes it and results carry both identities. + """ with logfire.span( "search.hydrate_results", domain="search", @@ -391,6 +427,12 @@ async def to_search_results( if temporal_by_source else None ), + project_id=result.project_id, + project_external_id=( + project_external_ids.get(result.project_id) + if project_external_ids + else None + ), ) ) return search_results diff --git a/src/basic_memory/repository/search_reader.py b/src/basic_memory/repository/search_reader.py index 0b5c81eb5..8d5d80b28 100644 --- a/src/basic_memory/repository/search_reader.py +++ b/src/basic_memory/repository/search_reader.py @@ -615,6 +615,9 @@ async def vector_only( ``candidate_limit`` is supplied only by a composed retrieval stage that already sized the shared candidate pool. """ + # An empty scope admits no rows; embedding the query would buy nothing. + if self.scope.is_empty: + return [] query_text = (query.search_text or "").strip() if candidate_limit is None: candidate_limit = self._candidate_limit(limit, offset, query_text) @@ -833,6 +836,8 @@ async def hybrid( ``max(vec, fts) + FUSION_BONUS * min(vec, fts)`` preserves the dominant signal and rewards dual-source agreement. """ + if self.scope.is_empty: + return [] query_text = (query.search_text or "").strip() rerank = self._active_rerank(query_text) query_start = time.perf_counter() diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index b21aca5b1..ca4eb0c14 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -16,14 +16,25 @@ from basic_memory.config import BasicMemoryConfig, DatabaseBackend from basic_memory.repository.embedding_provider_factory import create_embedding_provider from basic_memory.repository.rerank_provider_factory import create_rerank_provider +from basic_memory.repository.postgres_search_query import PostgresFts from basic_memory.repository.postgres_search_repository import PostgresSearchRepository +from basic_memory.repository.search_filters import FtsBackend from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.search_reader import ( + Reranking, + SearchReader, + SemanticSearch, + VectorRetrieval, +) from basic_memory.repository.search_repository_base import ChunkManifestRow, SearchIndexKey +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.search_trace import SearchTraceCollector from basic_memory.repository.semantic_vector_index_factory import ( create_semantic_vector_index, resolve_semantic_vector_index_name, + semantic_embedding_identity, ) +from basic_memory.repository.sqlite_search_query import SQLiteFts from basic_memory.runtime.vector_sync import VectorSyncBatchResult from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository from basic_memory.schemas.search import SearchItemType, SearchRetrievalMode @@ -265,8 +276,60 @@ def create_search_repository( ) +def create_search_reader( + session_maker: async_sessionmaker[AsyncSession], + scope: ProjectScope, + app_config: BasicMemoryConfig, + database_backend: Optional[DatabaseBackend] = None, +) -> SearchReader: + """Compose the read path over an explicit scope, with no project repository. + + Resolves the same shared embedding provider, vector adapter, and reranker that + ``create_search_repository`` hands a project repository, so a scoped search runs + the pipeline a project's own route runs, over a wider scope. Whether semantic + retrieval is available is decided here, once, from configuration; a semantic + query against a reader built without it fails as a disabled feature. + """ + backend = database_backend or app_config.database_backend + fts: FtsBackend = ( + PostgresFts(session_maker) + if backend == DatabaseBackend.POSTGRES + else SQLiteFts(session_maker) + ) + if not app_config.semantic_search_enabled: + return SearchReader(scope, fts) + + embedding_provider = create_embedding_provider(app_config) + vector_index_name, vector_index = create_semantic_vector_index( + session_maker=session_maker, + app_config=app_config, + database_backend=backend, + embedding_provider=embedding_provider, + ) + vector = VectorRetrieval( + index=vector_index, + index_name=vector_index_name, + embedding_provider=embedding_provider, + embedding_model=semantic_embedding_identity(embedding_provider), + vector_k=app_config.semantic_vector_k, + min_similarity=app_config.semantic_min_similarity, + ) + rerank_provider = create_rerank_provider(app_config) + rerank = ( + Reranking( + provider=rerank_provider, + candidates=app_config.reranker_candidates, + max_document_chars=app_config.reranker_max_document_chars, + ) + if rerank_provider is not None + else None + ) + return SearchReader(scope, fts, SemanticSearch(session_maker, scope, fts, vector, rerank)) + + __all__ = [ "SearchRepository", "SearchIndexRow", + "create_search_reader", "create_search_repository", ] diff --git a/src/basic_memory/repository/sqlite_vec_index.py b/src/basic_memory/repository/sqlite_vec_index.py index f85f40248..cfa45ddf7 100644 --- a/src/basic_memory/repository/sqlite_vec_index.py +++ b/src/basic_memory/repository/sqlite_vec_index.py @@ -79,7 +79,18 @@ async def _ensure_loaded(self, session: AsyncSession) -> None: "basic-memory under uv-managed or Homebrew Python, or disable " "semantic search." ) - await driver_connection.enable_load_extension(True) + try: + await driver_connection.enable_load_extension(True) + except AttributeError as exc: + # aiosqlite exposes the wrapper method even when the wrapped + # sqlite3.Connection was built without extension support, so + # calling it is the authoritative probe (#711). + raise SemanticDependenciesMissingError( + "This Python build does not support SQLite extension loading " + "(no enable_load_extension on sqlite3.Connection). Reinstall " + "basic-memory under uv-managed or Homebrew Python, or disable " + "semantic search." + ) from exc await driver_connection.load_extension(sqlite_vec.loadable_path()) await driver_connection.enable_load_extension(False) await session.execute(text("SELECT vec_version()")) diff --git a/src/basic_memory/schemas/search.py b/src/basic_memory/schemas/search.py index ba2ce0cc1..ddd1435a0 100644 --- a/src/basic_memory/schemas/search.py +++ b/src/basic_memory/schemas/search.py @@ -6,7 +6,7 @@ 3. Full-text search across content """ -from typing import Optional, List, Union, Any +from typing import Annotated, Optional, List, Union, Any from datetime import datetime from enum import Enum from pydantic import BaseModel, Field, ValidationInfo, field_validator, model_validator @@ -227,6 +227,17 @@ def has_boolean_operators(self) -> bool: return any(pattern in text for pattern in boolean_patterns) +class ScopedSearchQuery(SearchQuery): + """A search over an explicit set of projects in one database. + + ``project_ids`` are internal ids the caller has already authorized; the caller + decides what is visible and this route never widens it. The set is required and + may be empty, which answers no rows. There is no spelling for every project. + """ + + project_ids: list[Annotated[int, Field(strict=True, gt=0)]] + + class TemporalRangeValue(BaseModel): """One authored interval, as a caller sees it. @@ -299,6 +310,12 @@ class SearchResult(BaseModel): # of multiple kinds must not be a schema break later. temporal: Optional[List[TemporalResultMetadata]] = None + # The project this hit belongs to. The project route fills both from its path; + # the database-scoped route fills them from each row, since one page can span + # several projects. + project_id: Optional[int] = None + project_external_id: Optional[str] = None + class SearchResponse(BaseModel): """Wrapper for search results.""" diff --git a/src/basic_memory/services/scoped_search_service.py b/src/basic_memory/services/scoped_search_service.py new file mode 100644 index 000000000..38900b618 --- /dev/null +++ b/src/basic_memory/services/scoped_search_service.py @@ -0,0 +1,156 @@ +"""Search over an explicit set of projects in one database. + +The project route and this service run the same ``SearchReader``; only the scope +differs. Hydration stays inside the scope too: owning entities, relation endpoints, +project identities, and valid-time assertions are read only from the projects the +search was allowed to read, so a page never names something outside its scope. +""" + +from collections.abc import Iterable, Sequence + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker +from sqlalchemy.orm import load_only + +from basic_memory import db +from basic_memory.models import Entity, MemoryTimeIndex, Project +from basic_memory.repository.repository import SELECT_BY_IDS_CHUNK_SIZE +from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.search_reader import SearchReader +from basic_memory.repository.search_scope import ProjectScope +from basic_memory.schemas.search import SearchQuery +from basic_memory.services.search_service import ( + include_legacy_note_type_spellings, + prepare_search_query, + relaxed_fts_fallback_eligible, +) + + +def _chunks[T](values: Sequence[T]) -> Iterable[Sequence[T]]: + """Split a bind list at the shared per-statement parameter bound.""" + for start in range(0, len(values), SELECT_BY_IDS_CHUNK_SIZE): + yield values[start : start + SELECT_BY_IDS_CHUNK_SIZE] + + +class ScopedSearchService: + """Run and hydrate searches over one ``ProjectScope``.""" + + def __init__( + self, + session_maker: async_sessionmaker[AsyncSession], + scope: ProjectScope, + reader: SearchReader, + ) -> None: + self.session_maker = session_maker + self.scope = scope + self.reader = reader + + # --- Retrieval --- + + async def search( + self, + query: SearchQuery, + *, + limit: int, + offset: int, + ) -> list[SearchIndexRow]: + """One ranking over every project in scope.""" + prepared = prepare_search_query(query) + if prepared is None: + return [] + prepared = await include_legacy_note_type_spellings( + self.session_maker, self.scope, prepared + ) + allow_relaxed = relaxed_fts_fallback_eligible( + query, prepared.search_text, prepared.retrieval_mode + ) + return await self.reader.search( + prepared, limit=limit, offset=offset, allow_relaxed=allow_relaxed + ) + + async def count(self, query: SearchQuery) -> int: + """Exact full-text match count over every project in scope.""" + prepared = prepare_search_query(query) + if prepared is None: + return 0 + prepared = await include_legacy_note_type_spellings( + self.session_maker, self.scope, prepared + ) + allow_relaxed = relaxed_fts_fallback_eligible( + query, prepared.search_text, prepared.retrieval_mode + ) + return await self.reader.count(prepared, allow_relaxed=allow_relaxed) + + # --- Hydration, bounded by the scope --- + + async def get_entities_by_id(self, ids: Sequence[int]) -> Sequence[Entity]: + """The entities a page of hits refers to, read only from projects in scope. + + Only the fields result shaping reads are loaded. Ids reached through a scoped + search already belong to the scope; the predicate keeps that true by + construction rather than by trust. + """ + if not ids or self.scope.is_empty: + return [] + entities: list[Entity] = [] + async with db.scoped_session(self.session_maker) as session: + for chunk in _chunks(list(ids)): + result = await session.scalars( + select(Entity) + .where( + Entity.project_id.in_(self.scope.project_ids), + Entity.id.in_(chunk), + ) + .options( + load_only( + Entity.id, Entity.project_id, Entity.permalink, Entity.external_id + ) + ) + ) + entities.extend(result.all()) + return entities + + async def find_for_sources( + self, + session: AsyncSession, + sources: Iterable[tuple[str, int]], + ) -> Sequence[MemoryTimeIndex]: + """The valid-time assertions behind a page of hits, read only from projects in scope. + + Mirrors ``MemoryTimeIndexRepository.find_for_sources`` for a set of projects: + one statement per source type, chunked at the bind bound. + """ + if self.scope.is_empty: + return [] + ids_by_type: dict[str, list[int]] = {} + for source_type, source_id in sources: + ids_by_type.setdefault(source_type, []).append(source_id) + + rows: list[MemoryTimeIndex] = [] + for source_type, source_ids in ids_by_type.items(): + for chunk in _chunks(source_ids): + result = await session.scalars( + select(MemoryTimeIndex) + .where( + MemoryTimeIndex.project_id.in_(self.scope.project_ids), + MemoryTimeIndex.source_type == source_type, + MemoryTimeIndex.source_id.in_(chunk), + ) + .order_by(MemoryTimeIndex.source_id, MemoryTimeIndex.id) + ) + rows.extend(result.all()) + return rows + + async def project_external_ids( + self, + session: AsyncSession, + rows: Sequence[SearchIndexRow], + ) -> dict[int, str]: + """External ids for the projects a page of hits came from.""" + project_ids = sorted({row.project_id for row in rows}) + if not project_ids: + return {} + result = await session.execute( + select(Project.id, Project.external_id).where(Project.id.in_(project_ids)) + ) + return {int(project_id): str(external_id) for project_id, external_id in result.all()} diff --git a/src/basic_memory/services/search_service.py b/src/basic_memory/services/search_service.py index 0a0f70317..1ba056858 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -11,6 +11,7 @@ from dateparser import parse from fastapi import BackgroundTasks from loguru import logger +from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker import logfire @@ -24,6 +25,7 @@ SearchRepository, ) from basic_memory.repository.search_query import PreparedSearchQuery, relaxed_query_words +from basic_memory.repository.search_scope import ProjectScope from basic_memory.repository.search_trace import SearchTraceCollector from basic_memory.schemas.base import normalize_note_type from basic_memory.schemas.search import SearchQuery, SearchItemType, SearchRetrievalMode @@ -140,6 +142,132 @@ def _strip_nul(value: str) -> str: return value.replace("\x00", "") +def prepare_search_query(query: SearchQuery) -> PreparedSearchQuery | None: + """Normalize a ``SearchQuery`` into the prepared form every reader consumes. + + Returns ``None`` when the query names no criteria at all, so callers can answer + an empty page without touching storage. + """ + search_text = query.text + tags = query.tags + + # Support tag: shorthand by mapping to tags filter. + if search_text is not None: + search_text = search_text.strip() or None + if search_text and search_text.lower().startswith("tag:"): + tag_values = re.split(r"[,\s]+", search_text[4:].strip()) + parsed_tags = [t for t in tag_values if t] + if parsed_tags: + tags = parsed_tags + search_text = None + + after_date = ( + (query.after_date if isinstance(query.after_date, datetime) else parse(query.after_date)) + if query.after_date + else None + ) + + # Merge structured metadata filters (explicit + convenience fields). + metadata_filters: Optional[Dict[str, Any]] = None + if query.metadata_filters or tags or query.status: + metadata_filters = dict(query.metadata_filters or {}) + if tags: + metadata_filters.setdefault("tags", tags) + if query.status: + metadata_filters.setdefault("status", query.status) + + prepared = PreparedSearchQuery( + search_text=search_text, + permalink=query.permalink, + permalink_match=query.permalink_match, + title=query.title, + note_types=( + [normalize_note_type(note_type) for note_type in query.note_types] + if query.note_types + else None + ), + search_item_types=query.entity_types, + categories=query.categories, + after_date=after_date, + metadata_filters=metadata_filters, + file_path_prefix=query.file_path_prefix, + temporal=build_temporal_filter(query), + retrieval_mode=query.retrieval_mode or SearchRetrievalMode.FTS, + min_similarity=query.min_similarity, + ) + + has_criteria = bool( + prepared.search_text + or prepared.permalink + or prepared.permalink_match + or prepared.title + or prepared.note_types + or prepared.search_item_types + or prepared.categories + or prepared.after_date + or prepared.metadata_filters + # Normalized by SearchQuery, so only a real subtree reaches here. + or prepared.file_path_prefix + or prepared.temporal + ) + if not has_criteria: + logger.debug("no criteria passed to query") + return None + return prepared + + +async def include_legacy_note_type_spellings( + session_maker: async_sessionmaker[AsyncSession], + scope: ProjectScope, + prepared: PreparedSearchQuery, + *, + session: AsyncSession | None = None, +) -> PreparedSearchQuery: + """Expand canonical note-type filters to the exact legacy spellings stored in ``scope``. + + Search rows written before canonicalization preserve the owning entity's exact + type spelling. Including those spellings alongside the canonical values keeps an + upgrade searchable without requiring an eager full reindex. Only spellings from + projects in scope are read, so a scope cannot learn what another project stores. + """ + if not prepared.note_types: + return prepared + + canonical_note_types = set(prepared.note_types) + async with db.scoped_session(session_maker, session) as active_session: + stored_types = await active_session.scalars( + select(Entity.note_type).where(Entity.project_id.in_(scope.project_ids)).distinct() + ) + compatible_note_types = canonical_note_types | { + stored_type + for stored_type in stored_types.all() + if stored_type and normalize_note_type(stored_type) in canonical_note_types + } + return replace(prepared, note_types=sorted(compatible_note_types)) + + +def relaxed_fts_fallback_eligible( + query: SearchQuery, + search_text: str | None, + retrieval_mode: SearchRetrievalMode, +) -> bool: + """Whether a zero-result strict full-text query may retry with OR-joined terms.""" + if retrieval_mode != SearchRetrievalMode.FTS: + return False + if not search_text or not search_text.strip(): + return False + if '"' in search_text: + return False + if query.has_boolean_operators(): + return False + # Trigger: query has too few safe relaxed terms, explicit numeric identifiers, + # or only terms that would over-broaden under OR. + # Why: the shared helper preserves the old English guard while allowing + # whitespace-separated CJK terms that ASCII tokenization cannot see. + # Outcome: retry only when there is a backend-safe relaxed OR query. + return relaxed_query_words(search_text) is not None + + class SearchService: """Service for search operations. @@ -196,76 +324,7 @@ async def reindex_all(self, background_tasks: Optional[BackgroundTasks] = None) def prepare_query(self, query: SearchQuery) -> PreparedSearchQuery | None: """Normalize a SearchQuery into repository arguments.""" - search_text = query.text - tags = query.tags - - # Support tag: shorthand by mapping to tags filter. - if search_text is not None: - search_text = search_text.strip() or None - if search_text and search_text.lower().startswith("tag:"): - tag_values = re.split(r"[,\s]+", search_text[4:].strip()) - parsed_tags = [t for t in tag_values if t] - if parsed_tags: - tags = parsed_tags - search_text = None - - after_date = ( - ( - query.after_date - if isinstance(query.after_date, datetime) - else parse(query.after_date) - ) - if query.after_date - else None - ) - - # Merge structured metadata filters (explicit + convenience fields). - metadata_filters: Optional[Dict[str, Any]] = None - if query.metadata_filters or tags or query.status: - metadata_filters = dict(query.metadata_filters or {}) - if tags: - metadata_filters.setdefault("tags", tags) - if query.status: - metadata_filters.setdefault("status", query.status) - - prepared = PreparedSearchQuery( - search_text=search_text, - permalink=query.permalink, - permalink_match=query.permalink_match, - title=query.title, - note_types=( - [normalize_note_type(note_type) for note_type in query.note_types] - if query.note_types - else None - ), - search_item_types=query.entity_types, - categories=query.categories, - after_date=after_date, - metadata_filters=metadata_filters, - file_path_prefix=query.file_path_prefix, - temporal=build_temporal_filter(query), - retrieval_mode=query.retrieval_mode or SearchRetrievalMode.FTS, - min_similarity=query.min_similarity, - ) - - has_criteria = bool( - prepared.search_text - or prepared.permalink - or prepared.permalink_match - or prepared.title - or prepared.note_types - or prepared.search_item_types - or prepared.categories - or prepared.after_date - or prepared.metadata_filters - # Normalized by SearchQuery, so only a real subtree reaches here. - or prepared.file_path_prefix - or prepared.temporal - ) - if not has_criteria: - logger.debug("no criteria passed to query") - return None - return prepared + return prepare_search_query(query) @staticmethod def _prepared_has_filters(prepared: PreparedSearchQuery) -> bool: @@ -286,27 +345,12 @@ async def _include_legacy_note_type_spellings( session: AsyncSession | None = None, ) -> PreparedSearchQuery: """Expand canonical note-type filters to exact legacy entity spellings.""" - if not prepared.note_types: - return prepared - - canonical_note_types = set(prepared.note_types) - async with db.scoped_session(self.session_maker, session) as active_session: - stored_types_query = self.entity_repository.select(Entity.note_type).distinct() - stored_types_result = await self.entity_repository.execute_query( - active_session, - stored_types_query, - use_query_options=False, - ) - - # Search rows written before canonicalization preserve the owning entity's - # exact type spelling. Include those spellings alongside canonical values - # so an upgrade remains searchable without requiring an eager full reindex. - compatible_note_types = canonical_note_types | { - stored_type - for stored_type in stored_types_result.scalars().all() - if stored_type and normalize_note_type(stored_type) in canonical_note_types - } - return replace(prepared, note_types=sorted(compatible_note_types)) + return await include_legacy_note_type_spellings( + self.session_maker, + ProjectScope.single(self.repository.project_id), + prepared, + session=session, + ) async def _search_repository( self, @@ -482,20 +526,7 @@ def _is_relaxed_fts_fallback_eligible( retrieval_mode: SearchRetrievalMode, ) -> bool: """Check whether we should run relaxed OR fallback after strict FTS returns empty.""" - if retrieval_mode != SearchRetrievalMode.FTS: - return False - if not search_text or not search_text.strip(): - return False - if '"' in search_text: - return False - if query.has_boolean_operators(): - return False - # Trigger: query has too few safe relaxed terms, explicit numeric identifiers, - # or only terms that would over-broaden under OR. - # Why: the shared helper preserves the old English guard while allowing - # whitespace-separated CJK terms that ASCII tokenization cannot see. - # Outcome: retry only when there is a backend-safe relaxed OR query. - return relaxed_query_words(search_text) is not None + return relaxed_fts_fallback_eligible(query, search_text, retrieval_mode) @staticmethod def _generate_variants(text: str) -> Set[str]: diff --git a/tests/api/v2/test_scoped_search_router.py b/tests/api/v2/test_scoped_search_router.py new file mode 100644 index 000000000..41a8d3668 --- /dev/null +++ b/tests/api/v2/test_scoped_search_router.py @@ -0,0 +1,657 @@ +"""One query over an explicit set of projects, through the reader and the route. + +Only the embedding provider is a double; projects, entities, search rows, vectors, +retrieval, and hydration are real on whichever backend the session is configured +for. Search row ids are database-wide primary keys, so the corpus gives every row a +distinct id the way real data does. +""" + +import re +from dataclasses import dataclass +from datetime import datetime, timezone +from math import sqrt +from typing import Any +from unittest.mock import AsyncMock + +import pytest +from httpx import AsyncClient +from sqlalchemy import event, text +from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker + +import basic_memory.repository.search_repository as search_repository_module +from basic_memory import db +from basic_memory.api.v2.utils import to_search_results +from basic_memory.config import BasicMemoryConfig, DatabaseBackend +from basic_memory.models import Entity, MemoryTimeIndex, Project +from basic_memory.repository.postgres_search_repository import PostgresSearchRepository +from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.search_repository import create_search_reader +from basic_memory.repository.search_scope import ProjectScope +from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository +from basic_memory.schemas.search import SearchQuery, SearchRetrievalMode +from basic_memory.services.scoped_search_service import ScopedSearchService +from basic_memory.services.search_service import ( + include_legacy_note_type_spellings, + prepare_search_query, +) + +MARKERS = ("nebula", "comet", "quasar") +# Distinct per project, as real observation and relation primary keys are. +OBSERVATION_IDS = (5001, 5002, 5003) +RELATION_IDS = (6001, 6002, 6003) +SEMANTIC_MODES = [SearchRetrievalMode.VECTOR, SearchRetrievalMode.HYBRID] + + +class MarkerEmbeddingProvider: + """Unit vectors from marker words: deterministic, and different texts rank differently.""" + + model_name = "marker" + dimensions = 4 + + def __init__(self) -> None: + self.query_calls = 0 + + @staticmethod + def _vectorize(text: str) -> list[float]: + words = set(re.findall(r"[a-z]+", text.lower())) + axes = [1.0 if marker in words else 0.0 for marker in MARKERS] + axes.append(0.0 if any(axes) else 1.0) + norm = sqrt(sum(axis * axis for axis in axes)) + return [axis / norm for axis in axes] + + async def embed_query(self, text: str) -> list[float]: + self.query_calls += 1 + return self._vectorize(text) + + async def embed_documents(self, texts: list[str]) -> list[list[float]]: + return [self._vectorize(text) for text in texts] + + def runtime_log_attrs(self) -> dict[str, Any]: + return {} + + +@dataclass +class Corpus: + session_maker: async_sessionmaker[AsyncSession] + engine: AsyncEngine + config: BasicMemoryConfig + projects: list[Project] + entities: list[Entity] + provider: MarkerEmbeddingProvider + + def ids(self, count: int) -> list[int]: + return [project.id for project in self.projects[:count]] + + def service(self, project_ids: list[int]) -> ScopedSearchService: + scope = ProjectScope.of(project_ids) + reader = create_search_reader(self.session_maker, scope, self.config) + return ScopedSearchService(self.session_maker, scope, reader) + + +@pytest.fixture +async def corpus( + engine_factory, app_config: BasicMemoryConfig, monkeypatch: pytest.MonkeyPatch +) -> Corpus: + """Three projects with one note each, all about a nebula, embedded for real. + + Project 0 matches the query vector exactly; projects 1 and 2 carry a second + marker so their vectors sit at cosine 0.707. Project 2 spells its note type the + legacy way, so note-type expansion has something to find. + """ + engine, session_maker = engine_factory + app_config.semantic_search_enabled = True + app_config.reranker_enabled = False + app_config.semantic_min_similarity = 0.0 + provider = MarkerEmbeddingProvider() + monkeypatch.setattr( + search_repository_module, "create_embedding_provider", lambda _config: provider + ) + repository_type = ( + PostgresSearchRepository + if app_config.database_backend == DatabaseBackend.POSTGRES + else SQLiteSearchRepository + ) + now = datetime(2026, 9, 1, tzinfo=timezone.utc) + projects: list[Project] = [] + entities: list[Entity] = [] + for index, flavor in enumerate(("", " comet", " quasar")): + note_type = "Note" if index == 2 else "note" + async with db.scoped_session(session_maker) as session: + project = Project( + name=f"scope-{index}", permalink=f"scope-{index}", path=f"/scope-{index}" + ) + session.add(project) + await session.flush() + entity = Entity( + project_id=project.id, + title="shared nebula", + note_type=note_type, + content_type="text/markdown", + permalink="notes/shared", + file_path="notes/shared.md", + entity_metadata={"status": "open" if index == 0 else "closed", "tags": ["scope"]}, + created_at=now, + updated_at=now, + ) + session.add(entity) + await session.commit() + projects.append(project) + entities.append(entity) + + writer = repository_type( + session_maker, project.id, app_config=app_config, embedding_provider=provider + ) + await writer.init_search_index() + content = f"shared nebula{flavor}" + rows = [ + SearchIndexRow( + project_id=project.id, + id=entity.id, + entity_id=entity.id, + type="entity", + file_path=entity.file_path, + permalink=entity.permalink, + title=entity.title, + content_stems=content, + content_snippet=content, + metadata={"note_type": note_type}, + created_at=now, + updated_at=now, + ) + ] + for row_type, row_id in ( + ("observation", OBSERVATION_IDS[index]), + ("relation", RELATION_IDS[index]), + ): + rows.append( + SearchIndexRow( + project_id=project.id, + id=row_id, + entity_id=entity.id, + type=row_type, + file_path=entity.file_path, + permalink=f"notes/shared/{row_type}", + title="shared nebula", + content_stems=content, + content_snippet=f"{content} project {index} {row_type}", + category="fact" if row_type == "observation" else None, + from_id=entity.id if row_type == "relation" else None, + relation_type="relates_to" if row_type == "relation" else None, + created_at=now, + updated_at=now, + ) + ) + await writer.bulk_index_items(rows) + await writer.sync_entity_vectors(entity.id) + return Corpus(session_maker, engine, app_config, projects, entities, provider) + + +async def _mark_pending(corpus: Corpus, project_ids: list[int] | None = None) -> None: + """Take vectors out of play so a hybrid answer proves the lexical channel.""" + async with db.scoped_session(corpus.session_maker) as session: + if project_ids is None: + await session.execute( + text("UPDATE search_vector_chunks SET embedding_status = 'pending'") + ) + else: + for project_id in project_ids: + await session.execute( + text( + "UPDATE search_vector_chunks SET embedding_status = 'pending' " + "WHERE project_id = :project_id" + ), + {"project_id": project_id}, + ) + await session.commit() + + +async def _add_temporal_assertions(corpus: Corpus, project_ids: list[int]) -> None: + """Project 0's observation is current; every other project's ended in 2025.""" + async with db.scoped_session(corpus.session_maker) as session: + for index, (project, entity) in enumerate(zip(corpus.projects, corpus.entities)): + if project.id not in project_ids: + continue + current = index == 0 + session.add( + MemoryTimeIndex( + project_id=project.id, + entity_id=entity.id, + source_type="observation", + source_id=OBSERVATION_IDS[index], + time_kind="effective", + range_axis="date", + lower_value="2026-09-01" if current else "2025-01-01", + upper_value=None if current else "2025-02-01", + lower_inclusive=True, + upper_inclusive=False, + is_empty=False, + extractor="test", + source_text=( + "@effective[2026-09-01,)" + if current + else "@effective[2025-01-01,2025-02-01)" + ), + ) + ) + await session.commit() + + +# --- The reader over a scope --- + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", list(SearchRetrievalMode)) +@pytest.mark.parametrize("scope_size", [0, 1, 2]) +async def test_one_reader_answers_the_whole_scope( + corpus: Corpus, mode: SearchRetrievalMode, scope_size: int +) -> None: + """Every project in scope answers from one pipeline, and the pipeline only reads.""" + service = corpus.service(corpus.ids(scope_size)) + query = SearchQuery(text="nebula", retrieval_mode=mode) + statements: list[str] = [] + + def record(_conn, _cursor, statement, _params, _context, _many): + statements.append(statement) + + event.listen(corpus.engine.sync_engine, "before_cursor_execute", record) + try: + rows = await service.search(query, limit=100, offset=0) + finally: + event.remove(corpus.engine.sync_engine, "before_cursor_execute", record) + + assert len(rows) == scope_size * 3 + assert {row.project_id for row in rows} == set(corpus.ids(scope_size)) + assert len({(row.type, row.id) for row in rows}) == len(rows) + # One embedding per search, and none at all for an empty scope. + assert corpus.provider.query_calls == int(scope_size > 0 and mode != SearchRetrievalMode.FTS) + mutations = [ + sql + for sql in statements + if sql.lstrip().upper().startswith(("INSERT", "UPDATE", "DELETE", "DROP")) + ] + assert mutations == [] + if mode == SearchRetrievalMode.FTS: + assert await service.count(query) == scope_size * 3 + else: + assert all(row.matched_chunk_text for row in rows) + with pytest.raises(ValueError, match="Exact counts"): + await service.count(query) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", list(SearchRetrievalMode)) +async def test_pages_are_a_stable_slice_of_the_complete_order( + corpus: Corpus, mode: SearchRetrievalMode +) -> None: + service = corpus.service(corpus.ids(2)) + query = SearchQuery(text="nebula", retrieval_mode=mode) + + complete = await service.search(query, limit=100, offset=0) + pages = [ + row + for offset in range(0, len(complete), 2) + for row in await service.search(query, limit=2, offset=offset) + ] + + assert [(r.project_id, r.type, r.id, r.score) for r in pages] == [ + (r.project_id, r.type, r.id, r.score) for r in complete + ] + assert await service.search(query, limit=10, offset=100) == [] + + +@pytest.mark.asyncio +async def test_vectors_outside_the_current_manifest_are_not_served(corpus: Corpus) -> None: + """Pending, re-modelled, and stale-source chunks are invisible across the whole scope.""" + ids = corpus.ids(3) + async with db.scoped_session(corpus.session_maker) as session: + await session.execute( + text( + "UPDATE search_vector_chunks SET embedding_status = 'pending' WHERE project_id = :p" + ), + {"p": ids[0]}, + ) + await session.execute( + text( + "UPDATE search_vector_chunks SET embedding_model = 'other-model' WHERE project_id = :p" + ), + {"p": ids[1]}, + ) + await session.execute( + text("UPDATE search_vector_chunks SET source_hash = 'stale' WHERE project_id = :p"), + {"p": ids[2]}, + ) + await session.commit() + + query = SearchQuery(text="nebula", retrieval_mode=SearchRetrievalMode.VECTOR) + assert await corpus.service(ids).search(query, limit=100, offset=0) == [] + + +@pytest.mark.asyncio +async def test_hybrid_keeps_lexical_only_rows_below_the_threshold(corpus: Corpus) -> None: + ids = corpus.ids(2) + await _mark_pending(corpus, [ids[1]]) + query = SearchQuery(text="nebula", retrieval_mode=SearchRetrievalMode.HYBRID, min_similarity=1) + + rows = await corpus.service(ids).search(query, limit=100, offset=0) + + assert len(rows) == 6 + assert all(row.matched_chunk_text is None for row in rows if row.project_id == ids[1]) + assert max(row.score or 0.0 for row in rows) == pytest.approx(1.3) + assert all(1.0 <= (row.score or 0.0) <= 1.3 + 1e-6 for row in rows if row.project_id == ids[0]) + assert all(0.0 <= (row.score or 0.0) <= 1.0 for row in rows if row.project_id == ids[1]) + + +@pytest.mark.asyncio +async def test_legacy_note_type_spellings_expand_only_within_scope(corpus: Corpus) -> None: + """A scope learns the spellings its own projects store, and nothing about others.""" + prepared = prepare_search_query(SearchQuery(text="nebula", note_types=["note"])) + assert prepared is not None + canonical_only = ProjectScope.of(corpus.ids(2)) + legacy_project = ProjectScope.single(corpus.projects[2].id) + + within_canonical = await include_legacy_note_type_spellings( + corpus.session_maker, canonical_only, prepared + ) + within_legacy = await include_legacy_note_type_spellings( + corpus.session_maker, legacy_project, prepared + ) + + assert within_canonical.note_types == ["note"] + assert within_legacy.note_types == ["Note", "note"] + # And the expanded filter finds the legacy rows when that project is in scope. + query = SearchQuery(text="nebula", note_types=["note"]) + assert len(await corpus.service(corpus.ids(3)).search(query, limit=100, offset=0)) == 9 + assert ( + len(await corpus.service([corpus.projects[2].id]).search(query, limit=100, offset=0)) == 3 + ) + + +@pytest.mark.asyncio +async def test_hydration_reads_only_projects_in_scope(corpus: Corpus) -> None: + ids = corpus.ids(2) + await _add_temporal_assertions(corpus, ids) + service = corpus.service([ids[0]]) + all_entity_ids = [entity.id for entity in corpus.entities] + + entities = await service.get_entities_by_id(all_entity_ids) + # Search before opening the hydration session: the test pool holds one connection. + rows = await service.search(SearchQuery(text="nebula"), limit=100, offset=0) + async with db.scoped_session(corpus.session_maker) as session: + assertions = await service.find_for_sources( + session, [("observation", source_id) for source_id in OBSERVATION_IDS] + ) + external_ids = await service.project_external_ids(session, rows) + + assert [entity.id for entity in entities] == [corpus.entities[0].id] + assert [row.project_id for row in assertions] == [ids[0]] + assert external_ids == {ids[0]: corpus.projects[0].external_id} + # An empty scope hydrates nothing, and an empty page names no projects. + unscoped = corpus.service([]) + assert await unscoped.get_entities_by_id(all_entity_ids) == [] + async with db.scoped_session(corpus.session_maker) as session: + assert await unscoped.find_for_sources(session, [("observation", OBSERVATION_IDS[0])]) == [] + assert await service.project_external_ids(session, []) == {} + + +# --- The route --- + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", list(SearchRetrievalMode)) +async def test_filters_narrow_the_scope_and_results_carry_project_identity( + corpus: Corpus, client: AsyncClient, mode: SearchRetrievalMode +) -> None: + ids = corpus.ids(2) + + response = await client.request( + "QUERY", + "/v2/search/", + json={ + "project_ids": ids, + "text": "nebula", + "retrieval_mode": mode.value, + "metadata_filters": {"status": "open"}, + "note_types": ["note"], + "file_path_prefix": "notes", + "entity_types": ["observation"], + "categories": ["fact"], + }, + ) + + assert response.status_code == 200, response.text + data = response.json() + assert len(data["results"]) == 1 + row = data["results"][0] + assert row["project_id"] == ids[0] + assert row["project_external_id"] == corpus.projects[0].external_id + assert row["external_id"] == corpus.entities[0].external_id + assert row["observation_id"] == OBSERVATION_IDS[0] + assert row["entity"] == "notes/shared" + assert data["total_is_exact"] == (mode == SearchRetrievalMode.FTS) + assert data["total"] == int(mode == SearchRetrievalMode.FTS) + assert not data["has_more"] + assert response.headers["cache-control"] == "no-store" + assert response.headers["accept-query"] == "application/json" + + +@pytest.mark.asyncio +async def test_scope_shapes_the_answer(corpus: Corpus, client: AsyncClient) -> None: + for ids, count in [(corpus.ids(3), 9), ([corpus.projects[1].id], 3), ([], 0)]: + response = await client.request( + "QUERY", "/v2/search/", json={"project_ids": ids, "text": "nebula"} + ) + assert response.status_code == 200, response.text + assert response.json()["total"] == count + assert {r["project_id"] for r in response.json()["results"]} <= set(ids) + + +@pytest.mark.asyncio +async def test_post_is_the_documented_twin_of_query( + corpus: Corpus, client: AsyncClient, app +) -> None: + body = {"project_ids": corpus.ids(2), "text": "nebula"} + + posted = await client.post("/v2/search/", json=body) + queried = await client.request("QUERY", "/v2/search/", json=body) + + assert posted.status_code == 200, posted.text + assert posted.json() == queried.json() + assert set(app.openapi()["paths"]["/v2/search/"]) == {"post"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("scope", [None, [0], [True], ["1"]]) +async def test_route_rejects_an_unspecified_or_invalid_scope( + client: AsyncClient, scope: object +) -> None: + response = await client.request( + "QUERY", "/v2/search/", json={"text": "nebula", "project_ids": scope} + ) + assert response.status_code == 422 + response = await client.request("QUERY", "/v2/search/", json={"text": "nebula"}) + assert response.status_code == 422 + response = await client.request( + "QUERY", "/v2/search/", params={"page": 0}, json={"text": "nebula", "project_ids": [1]} + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", list(SearchRetrievalMode)) +async def test_temporal_filter_narrows_and_explains_within_scope( + corpus: Corpus, client: AsyncClient, mode: SearchRetrievalMode +) -> None: + ids = corpus.ids(2) + await _add_temporal_assertions(corpus, ids) + + response = await client.request( + "QUERY", + "/v2/search/", + json={ + "project_ids": ids, + "text": "nebula", + "retrieval_mode": mode.value, + "valid_at": "2026-09-14", + "note_types": ["note"], + }, + ) + + assert response.status_code == 200, response.text + payload = response.json() + assert payload["temporal_applied"] is True + assert len(payload["results"]) == 1 + assert payload["results"][0]["project_id"] == ids[0] + assert payload["results"][0]["temporal"][0]["source_text"] == "@effective[2026-09-01,)" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", list(SearchRetrievalMode)) +async def test_api_pagination_and_empty_later_page( + corpus: Corpus, client: AsyncClient, mode: SearchRetrievalMode +) -> None: + query = {"project_ids": corpus.ids(2), "text": "nebula", "retrieval_mode": mode.value} + for page, count, has_more in [(1, 2, True), (3, 2, False), (4, 0, False)]: + response = await client.request( + "QUERY", "/v2/search/", params={"page": page, "page_size": 2}, json=query + ) + assert response.status_code == 200, response.text + assert len(response.json()["results"]) == count + assert response.json()["has_more"] == has_more + + +@pytest.mark.asyncio +async def test_semantic_modes_need_the_semantic_stack(corpus: Corpus, client: AsyncClient) -> None: + corpus.config.semantic_search_enabled = False + scope = {"project_ids": [corpus.projects[0].id]} + + for mode in SEMANTIC_MODES: + response = await client.request( + "QUERY", "/v2/search/", json={**scope, "text": "nebula", "retrieval_mode": mode.value} + ) + assert response.status_code == 400, response.text + assert "disabled" in response.json()["detail"] + assert corpus.provider.query_calls == 0 + + no_criteria = await client.request("QUERY", "/v2/search/", json=scope) + assert no_criteria.status_code == 200 and not no_criteria.json()["results"] + lexical = await client.request("QUERY", "/v2/search/", json={**scope, "text": "nebula"}) + assert lexical.status_code == 200, lexical.text + assert lexical.json()["total"] == 3 + + +@pytest.mark.asyncio +async def test_project_route_results_carry_project_identity( + corpus: Corpus, client: AsyncClient +) -> None: + project = corpus.projects[0] + + response = await client.post( + f"/v2/projects/{project.external_id}/search/", json={"text": "nebula"} + ) + + assert response.status_code == 200, response.text + results = response.json()["results"] + assert len(results) == 3 + assert {row["project_id"] for row in results} == {project.id} + assert {row["project_external_id"] for row in results} == {project.external_id} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["fts", "hybrid"]) +@pytest.mark.parametrize("query", ["Did nebula go hiking at sunrise?", "foo None: + ids = corpus.ids(2) + # Remove the vector channel so a hybrid success proves lexical recovery. + await _mark_pending(corpus) + body = {"project_ids": ids, "text": query, "retrieval_mode": mode} + + complete = await client.request("QUERY", "/v2/search/", json=body) + + assert complete.status_code == 200, complete.text + expected = complete.json()["results"] + assert len(expected) == 6 + assert {row["project_id"] for row in expected} == set(ids) + if mode == "fts": + assert complete.json()["total"] == 6 + pages = [] + for page in range(1, 5): + response = await client.request( + "QUERY", "/v2/search/", json=body, params={"page": page, "page_size": 2} + ) + assert response.status_code == 200, response.text + pages.extend(response.json()["results"]) + assert response.json()["has_more"] == (page < 3) + assert pages == expected + assert corpus.provider.query_calls == (5 if mode == "hybrid" else 0) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["fts", "hybrid"]) +async def test_relaxation_keeps_strict_matches_on_deep_pages( + corpus: Corpus, client: AsyncClient, mode: str +) -> None: + # Three terms permit relaxation, but strict matches anywhere in the selected + # scope must suppress it even after the final page. + await _mark_pending(corpus) + body = { + "project_ids": corpus.ids(2), + "text": "shared nebula observation", + "retrieval_mode": mode, + "min_similarity": 1, + } + + first = await client.request("QUERY", "/v2/search/", json=body) + later = await client.request("QUERY", "/v2/search/", json=body, params={"page": 2}) + + assert first.status_code == 200, first.text + assert len(first.json()["results"]) == 2 + assert later.status_code == 200, later.text + assert later.json()["results"] == [] + if mode == "fts": + assert later.json()["total"] == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("query", ["foo None: + response = await client.request( + "QUERY", + "/v2/search/", + json={"project_ids": [corpus.projects[0].id], "text": query, "retrieval_mode": mode}, + ) + + assert response.status_code == 200, response.text + assert len(response.json()["results"]) == (3 if mode == "hybrid" else 0) + if mode == "fts": + assert response.json()["total"] == 0 + assert corpus.provider.query_calls == int(mode == "hybrid") + + +# --- Hydration shaping --- + + +@pytest.mark.asyncio +async def test_results_name_their_project_when_the_caller_knows_it() -> None: + now = datetime(2026, 9, 1, tzinfo=timezone.utc) + row = SearchIndexRow( + project_id=7, + id=1, + type="entity", + title="Alpha", + file_path="alpha.md", + created_at=now, + updated_at=now, + ) + lookup = AsyncMock() + lookup.get_entities_by_id = AsyncMock(return_value=[]) + + named = await to_search_results(lookup, [row], project_external_ids={7: "project-7"}) + unnamed = await to_search_results(lookup, [row]) + + assert (named[0].project_id, named[0].project_external_id) == (7, "project-7") + assert (unnamed[0].project_id, unnamed[0].project_external_id) == (7, None) diff --git a/tests/api/v2/test_search_router_telemetry.py b/tests/api/v2/test_search_router_telemetry.py index d1eec98ba..eeb459975 100644 --- a/tests/api/v2/test_search_router_telemetry.py +++ b/tests/api/v2/test_search_router_telemetry.py @@ -32,7 +32,9 @@ def fake_span(name: str, **attrs): operations.append((name, attrs)) yield - async def fake_to_search_results(entity_service, results, *, temporal_by_source=None): + async def fake_to_search_results( + entity_service, results, *, temporal_by_source=None, project_external_ids=None + ): return [] monkeypatch.setattr(logfire, "span", fake_span) @@ -48,6 +50,7 @@ async def fake_to_search_results(entity_service, results, *, temporal_by_source= session_maker=object(), read_cache=None, response=http_response, + internal_project_id=1, project_id="11111111-1111-1111-1111-111111111111", page=2, page_size=5, diff --git a/tests/repository/test_search_reader.py b/tests/repository/test_search_reader.py index cccf2668b..de29903cb 100644 --- a/tests/repository/test_search_reader.py +++ b/tests/repository/test_search_reader.py @@ -281,3 +281,20 @@ async def test_hybrid_trace_records_the_stable_pool_refetch(): assert trace.stable_pool_refetched is True assert [row.id for row in page] == [4] + + +@pytest.mark.asyncio +async def test_empty_scope_answers_nothing_without_embedding_or_lexical_work() -> None: + """An empty scope admits no rows, so neither leg runs and the query is never embedded.""" + fts = FakeFts([FakeRow(id=1)]) + embed_query = AsyncMock(side_effect=AssertionError("must not embed for an empty scope")) + semantic = SemanticSearch( + cast(Any, None), ProjectScope.of([]), fts, fake_vector_retrieval(embed_query=embed_query) + ) + vector_query = PreparedSearchQuery( + search_text="auth", retrieval_mode=SearchRetrievalMode.VECTOR + ) + + assert await semantic.vector_only(vector_query, limit=10, offset=0) == [] + assert await semantic.hybrid(HYBRID_QUERY, limit=10, offset=0) == [] + assert fts.queries == [] diff --git a/tests/repository/test_search_reader_factory.py b/tests/repository/test_search_reader_factory.py new file mode 100644 index 000000000..88c96de26 --- /dev/null +++ b/tests/repository/test_search_reader_factory.py @@ -0,0 +1,115 @@ +"""Composing a SearchReader over a scope without a project repository.""" + +from typing import Any +from unittest.mock import MagicMock + +import pytest + +import basic_memory.repository.search_repository as search_repository_module +from basic_memory.config import BasicMemoryConfig, DatabaseBackend +from basic_memory.repository.postgres_search_query import PostgresFts +from basic_memory.repository.search_reader import Reranking +from basic_memory.repository.search_repository import create_search_reader +from basic_memory.repository.search_scope import ProjectScope +from basic_memory.repository.semantic_vector_index_factory import semantic_embedding_identity +from basic_memory.repository.sqlite_search_query import SQLiteFts + +SCOPE = ProjectScope.of([3, 1]) + + +class _StubEmbeddingProvider: + model_name = "stub" + dimensions = 4 + + async def embed_query(self, text: str) -> list[float]: + return [1.0, 0.0, 0.0, 0.0] + + async def embed_documents(self, texts: list[str]) -> list[list[float]]: + return [[1.0, 0.0, 0.0, 0.0] for _ in texts] + + def runtime_log_attrs(self) -> dict[str, Any]: + return {} + + +class _StubReranker: + model_name = "stub-reranker" + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + return [0.5 for _ in documents] + + def runtime_log_attrs(self) -> dict[str, Any]: + return {} + + +def _config(backend: DatabaseBackend, **overrides: object) -> BasicMemoryConfig: + return BasicMemoryConfig( + env="test", + projects={"test-project": "/tmp/test"}, + default_project="test-project", + database_backend=backend, + **overrides, + ) + + +@pytest.mark.parametrize( + ("backend", "fts_type"), + [(DatabaseBackend.SQLITE, SQLiteFts), (DatabaseBackend.POSTGRES, PostgresFts)], +) +def test_reader_without_semantic_search_is_full_text_only(monkeypatch, backend, fts_type): + """Disabled semantic search never resolves a provider, and the engine picks the backend.""" + monkeypatch.setattr( + search_repository_module, + "create_embedding_provider", + lambda _config: pytest.fail("a full-text reader must not load an embedding provider"), + ) + + reader = create_search_reader( + MagicMock(), SCOPE, _config(backend, semantic_search_enabled=False) + ) + + assert reader.scope == SCOPE + assert reader.semantic is None + assert isinstance(reader.fts, fts_type) + + +@pytest.mark.parametrize("reranker", [None, _StubReranker()]) +def test_reader_with_semantic_search_composes_the_shared_stack(monkeypatch, reranker): + """The reader gets the same provider, adapter, and reranker a project repository would.""" + provider = _StubEmbeddingProvider() + index = MagicMock() + captured: dict[str, Any] = {} + + def fake_create_index(**kwargs: Any) -> tuple[str, Any]: + captured.update(kwargs) + return "sqlite-vec", index + + monkeypatch.setattr(search_repository_module, "create_embedding_provider", lambda _c: provider) + monkeypatch.setattr(search_repository_module, "create_semantic_vector_index", fake_create_index) + monkeypatch.setattr(search_repository_module, "create_rerank_provider", lambda _c: reranker) + config = _config( + DatabaseBackend.SQLITE, + semantic_search_enabled=True, + semantic_vector_k=7, + semantic_min_similarity=0.25, + reranker_candidates=9, + reranker_max_document_chars=123, + ) + + reader = create_search_reader(MagicMock(), SCOPE, config) + + semantic = reader.semantic + assert semantic is not None + assert semantic.scope == SCOPE and semantic.fts is reader.fts + assert semantic.vector.index is index + assert semantic.vector.index_name == "sqlite-vec" + assert semantic.vector.embedding_provider is provider + assert semantic.vector.embedding_model == semantic_embedding_identity(provider) + assert (semantic.vector.vector_k, semantic.vector.min_similarity) == (7, 0.25) + # The adapter is built for the database, not for a project. + assert captured["database_backend"] == DatabaseBackend.SQLITE + assert captured["embedding_provider"] is provider + assert "project_id" not in captured + if reranker is None: + assert semantic.rerank is None + else: + assert semantic.rerank == Reranking(provider=reranker, candidates=9, max_document_chars=123) diff --git a/tests/repository/test_sqlite_vec_index.py b/tests/repository/test_sqlite_vec_index.py new file mode 100644 index 000000000..b5c613a5d --- /dev/null +++ b/tests/repository/test_sqlite_vec_index.py @@ -0,0 +1,41 @@ +"""sqlite-vec adapter behavior that does not need a database.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest +from sqlalchemy.exc import OperationalError as SAOperationalError + +from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError +from basic_memory.repository.semantic_vector_index import VectorIndexScope +from basic_memory.repository.sqlite_vec_index import SQLiteVecIndex + + +@pytest.mark.asyncio +async def test_missing_extension_support_is_a_dependency_error() -> None: + """A Python whose sqlite3 cannot load extensions gets the keyword-only fallback error. + + aiosqlite exposes ``enable_load_extension`` even when the wrapped connection does + not, so the attribute probe passes and only the call reveals it (#711). The adapter + must raise the same typed error the repository does, or a scoped vector search on + such a host answers 500 instead of an actionable 400. + """ + index = SQLiteVecIndex( + MagicMock(), + VectorIndexScope(namespace="test", embedding_identity="test", dimensions=4), + ) + driver = SimpleNamespace( + enable_load_extension=AsyncMock(side_effect=AttributeError("enable_load_extension")) + ) + connection = AsyncMock() + connection.get_raw_connection = AsyncMock( + return_value=SimpleNamespace(driver_connection=driver) + ) + session = AsyncMock() + session.execute = AsyncMock( + side_effect=SAOperationalError("SELECT vec_version()", {}, Exception("no such function")) + ) + session.connection = AsyncMock(return_value=connection) + + with pytest.raises(SemanticDependenciesMissingError, match="extension loading"): + await index._ensure_loaded(session)