From bc9ca0744ffe4296d7d597b4dd9b7c73c2d63f3f Mon Sep 17 00:00:00 2001 From: phernandez Date: Tue, 18 Feb 2025 21:29:11 -0600 Subject: [PATCH] fix: search query pagination params --- src/basic_memory/api/routers/resource_router.py | 14 +++++++------- src/basic_memory/mcp/tools/search.py | 4 +--- 2 files changed, 8 insertions(+), 10 deletions(-) diff --git a/src/basic_memory/api/routers/resource_router.py b/src/basic_memory/api/routers/resource_router.py index 3bc1c4bb..0feeae76 100644 --- a/src/basic_memory/api/routers/resource_router.py +++ b/src/basic_memory/api/routers/resource_router.py @@ -21,16 +21,16 @@ from basic_memory.schemas.search import SearchQuery, SearchItemType router = APIRouter(prefix="/resource", tags=["resources"]) -def get_entity_ids(item: SearchIndexRow) -> list[int]: +def get_entity_ids(item: SearchIndexRow) -> set[int]: match item.type: case SearchItemType.ENTITY: - return [item.id] + return {item.id} case SearchItemType.OBSERVATION: - return [item.entity_id] # pyright: ignore [reportReturnType] + return {item.entity_id} # pyright: ignore [reportReturnType] case SearchItemType.RELATION: from_entity = item.from_id to_entity = item.to_id # pyright: ignore [reportReturnType] - return [from_entity, to_entity] if to_entity else [from_entity] # pyright: ignore [reportReturnType] + return {from_entity, to_entity} if to_entity else {from_entity} # pyright: ignore [reportReturnType] case _: # pragma: no cover raise ValueError(f"Unexpected type: {item.type}") @@ -70,9 +70,9 @@ async def get_resource_content( if not search_results: raise HTTPException(status_code=404, detail=f"Resource not found: {identifier}") - # get the entities related to the search results - entity_ids = [id for result in search_results for id in get_entity_ids(result)] - results = await entity_service.get_entities_by_id(entity_ids) + # get the deduplicated entities related to the search results + entity_ids = {id for result in search_results for id in get_entity_ids(result)} + results = await entity_service.get_entities_by_id(list(entity_ids)) # return single response if len(results) == 1: diff --git a/src/basic_memory/mcp/tools/search.py b/src/basic_memory/mcp/tools/search.py index b3445db7..95259018 100644 --- a/src/basic_memory/mcp/tools/search.py +++ b/src/basic_memory/mcp/tools/search.py @@ -29,7 +29,5 @@ async def search(query: SearchQuery, page: int = 1, page_size: int = 10) -> Sear """ with logfire.span("Searching for {query}", query=query): # pyright: ignore [reportGeneralTypeIssues] logger.info(f"Searching for {query}") - response = await call_post( - client, f"/search/?page={page}&page_size={page_size}", json=query.model_dump() - ) + response = await call_post(client, "/search/", json=query.model_dump(), params={"page": page, "page_size": page_size}) return SearchResponse.model_validate(response.json())