|
| 1 | +"""Real API/MCP pagination over the same single-pass semantic repository.""" |
| 2 | + |
| 3 | +import json |
| 4 | + |
| 5 | +import pytest |
| 6 | +from fastapi import FastAPI |
| 7 | +from fastmcp import Client, FastMCP |
| 8 | +from httpx import AsyncClient |
| 9 | +from sqlalchemy import func, update |
| 10 | +from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker |
| 11 | + |
| 12 | +from basic_memory import db |
| 13 | + |
| 14 | +from basic_memory.deps.services import get_search_service_v2_external |
| 15 | +from basic_memory.models import Project |
| 16 | +from basic_memory.services.search_service import SearchService |
| 17 | +from tests.repository.test_rerank_pipeline import ( |
| 18 | + BackendSearchRepository, |
| 19 | + _FakeReranker, |
| 20 | + rerank_search_repository as rerank_search_repository, |
| 21 | +) |
| 22 | +from tests.repository.test_stable_rerank_pagination import ( |
| 23 | + pagination_repository as pagination_repository, |
| 24 | +) |
| 25 | + |
| 26 | + |
| 27 | +@pytest.fixture |
| 28 | +def session_maker( |
| 29 | + engine_factory: tuple[AsyncEngine, async_sessionmaker[AsyncSession]], |
| 30 | +) -> async_sessionmaker[AsyncSession]: |
| 31 | + return engine_factory[1] |
| 32 | + |
| 33 | + |
| 34 | +@pytest.mark.asyncio |
| 35 | +@pytest.mark.parametrize("mode", ["vector", "hybrid"]) |
| 36 | +async def test_api_and_mcp_keep_probe_and_page_boundaries( |
| 37 | + pagination_repository: BackendSearchRepository, |
| 38 | + search_service: SearchService, |
| 39 | + app: FastAPI, |
| 40 | + client: AsyncClient, |
| 41 | + test_project: Project, |
| 42 | + mcp_server: FastMCP, |
| 43 | + mode: str, |
| 44 | +) -> None: |
| 45 | + v2_project_url = f"/v2/projects/{test_project.external_id}" |
| 46 | + repo = pagination_repository |
| 47 | + repo._rerank_provider = _FakeReranker({"Note 00": 0.8, "Note 01": 0.9}) |
| 48 | + async with db.scoped_session(repo.session_maker) as session: |
| 49 | + await session.execute( |
| 50 | + update(Project).where(Project.id == test_project.id).values(last_indexed_at=func.now()) |
| 51 | + ) |
| 52 | + await session.commit() |
| 53 | + search_service.repository = repo |
| 54 | + app.dependency_overrides[get_search_service_v2_external] = lambda: search_service |
| 55 | + query = {"text": "auth session token", "retrieval_mode": mode, "min_similarity": 0.5} |
| 56 | + response = await client.request( |
| 57 | + "QUERY", f"{v2_project_url}/search/", json=query, params={"page_size": 100} |
| 58 | + ) |
| 59 | + assert response.status_code == 200, response.text |
| 60 | + expected = response.json()["results"] |
| 61 | + pages = [] |
| 62 | + async with Client(mcp_server) as mcp: |
| 63 | + for page in range(1, (len(expected) + 4) // 5 + 2): |
| 64 | + result = await mcp.call_tool( |
| 65 | + "search_notes", |
| 66 | + { |
| 67 | + "project": test_project.name, |
| 68 | + "query": "auth session token", |
| 69 | + "search_type": mode, |
| 70 | + "min_similarity": 0.5, |
| 71 | + "page": page, |
| 72 | + "page_size": 5, |
| 73 | + "output_format": "json", |
| 74 | + }, |
| 75 | + ) |
| 76 | + payload = json.loads(result.content[0].text) |
| 77 | + assert payload["has_more"] == (page * 5 < len(expected)) |
| 78 | + assert payload["total_is_exact"] is False |
| 79 | + pages.extend(payload["results"]) |
| 80 | + assert [row["permalink"] for row in pages] == [row["permalink"] for row in expected] |
0 commit comments