diff --git a/src/basic_memory/api/v2/routers/search_router.py b/src/basic_memory/api/v2/routers/search_router.py index eefc82da..be80d4b4 100644 --- a/src/basic_memory/api/v2/routers/search_router.py +++ b/src/basic_memory/api/v2/routers/search_router.py @@ -14,7 +14,7 @@ from basic_memory.repository.semantic_errors import ( SemanticDependenciesMissingError, SemanticSearchDisabledError, ) -from basic_memory.schemas.search import SearchQuery, SearchResponse +from basic_memory.schemas.search import SearchQuery, SearchResponse, SearchRetrievalMode from basic_memory.deps import ( SearchServiceV2ExternalDep, EntityServiceV2ExternalDep, @@ -67,6 +67,7 @@ async def search( has_filters=bool(query.note_types or query.entity_types or query.metadata_filters), ): offset = (page - 1) * page_size + exact_count_available = query.retrieval_mode == SearchRetrievalMode.FTS try: with logfire.span( "api.search.search.execute_query", @@ -76,10 +77,14 @@ async def search( page=page, page_size=page_size, ): - results, total = await asyncio.gather( - search_service.search(query, limit=page_size, offset=offset), - search_service.count(query), - ) + if exact_count_available: + results, total = await asyncio.gather( + search_service.search(query, limit=page_size, offset=offset), + search_service.count(query), + ) + else: + results = await search_service.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: @@ -94,7 +99,15 @@ async def search( phase="paginate_results", result_count=len(results), ): - has_more = offset + len(results) < total + if exact_count_available: + has_more = offset + len(results) < total + else: + # Trigger: semantic modes would need another vector/hybrid retrieval to count. + # Why: search requests should not pay for a second semantic pass. + # Outcome: preserve probe pagination for semantic search and leave total at 0. + has_more = len(results) > page_size + if has_more: + results = results[:page_size] with logfire.span( "api.search.search.hydrate_results", diff --git a/src/basic_memory/repository/search_repository_base.py b/src/basic_memory/repository/search_repository_base.py index 6a7a0590..c5367295 100644 --- a/src/basic_memory/repository/search_repository_base.py +++ b/src/basic_memory/repository/search_repository_base.py @@ -260,25 +260,10 @@ class SearchRepositoryBase(ABC): retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, ) -> int: - """Count results for retrieval modes that cannot use a backend COUNT query.""" - # Trigger: vector and hybrid modes rank after embedding lookup, filtering, and fusion. - # Why: that scoring pipeline is not expressible as a portable database COUNT query. - # Outcome: fetch the bounded candidate set and count the final in-memory results. - results = await self.search( - search_text=search_text, - permalink=permalink, - permalink_match=permalink_match, - title=title, - note_types=note_types, - after_date=after_date, - search_item_types=search_item_types, - metadata_filters=metadata_filters, - retrieval_mode=retrieval_mode, - min_similarity=min_similarity, - limit=VECTOR_FILTER_SCAN_LIMIT, - offset=0, - ) - return len(results) + """Count results when a backend-specific COUNT query is available.""" + if retrieval_mode != SearchRetrievalMode.FTS: + raise ValueError("Exact counts are only supported for full-text search retrieval.") + raise NotImplementedError("Backend search repositories must implement full-text counts.") # ------------------------------------------------------------------ # Abstract methods — semantic search (backend-specific DB operations) diff --git a/tests/api/v2/test_search_router.py b/tests/api/v2/test_search_router.py index 58cc091a..a7f93ce2 100644 --- a/tests/api/v2/test_search_router.py +++ b/tests/api/v2/test_search_router.py @@ -456,6 +456,56 @@ async def test_search_router_returns_400_for_invalid_vector_query( assert "Vector retrieval requires a text query" in response.json()["detail"] +@pytest.mark.asyncio +async def test_semantic_search_uses_probe_pagination_without_count( + client: AsyncClient, + app, + v2_project_url: str, +): + """Semantic searches should not run an extra count query.""" + now = datetime.now(timezone.utc) + fake_rows = [ + SearchIndexRow( + project_id=1, + id=row_id, + type="entity", + file_path=f"notes/semantic-{row_id}.md", + created_at=now, + updated_at=now, + title=f"Semantic Result {row_id}", + permalink=f"notes/semantic-{row_id}", + score=1.0 - (row_id / 10), + ) + for row_id in range(1, 4) + ] + + class FakeSearchService: + async def search(self, query, *, limit, offset): + assert query.retrieval_mode.value == "vector" + assert limit == 3 + assert offset == 0 + return fake_rows + + async def count(self, *args, **kwargs): + raise AssertionError("semantic search must not run count") + + app.dependency_overrides[get_search_service_v2_external] = lambda: FakeSearchService() + try: + response = await client.post( + f"{v2_project_url}/search/", + json={"text": "semantic query", "retrieval_mode": "vector"}, + params={"page": 1, "page_size": 2}, + ) + finally: + app.dependency_overrides.pop(get_search_service_v2_external, None) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 0 + assert data["has_more"] is True + assert len(data["results"]) == 2 + + @pytest.mark.asyncio async def test_search_has_more_when_more_results_exist( client: AsyncClient, diff --git a/tests/repository/test_semantic_search_base.py b/tests/repository/test_semantic_search_base.py index 6998d199..ba4e08ef 100644 --- a/tests/repository/test_semantic_search_base.py +++ b/tests/repository/test_semantic_search_base.py @@ -318,6 +318,25 @@ async def test_sqlite_hybrid_search_raises_disabled_error(search_repository): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("retrieval_mode", [SearchRetrievalMode.VECTOR, SearchRetrievalMode.HYBRID]) +async def test_count_rejects_semantic_modes_without_running_search(monkeypatch, retrieval_mode): + """Semantic counts must not materialize vector or hybrid retrieval.""" + repo = _ConcreteRepo() + search_calls = [] + + async def fail_if_search_runs(**kwargs): + search_calls.append(kwargs) + return [] + + monkeypatch.setattr(repo, "search", fail_if_search_runs) + + with pytest.raises(ValueError, match="Exact counts are only supported for full-text search"): + await repo.count(search_text="semantic query", retrieval_mode=retrieval_mode) + + assert search_calls == [] + + @pytest.mark.asyncio async def test_sync_entity_vectors_batch_flushes_at_configured_threshold(monkeypatch): """Batch sync should flush queued jobs at semantic_embedding_sync_batch_size boundaries."""