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)