From 2bc35f34e37ec9a0dfc7350d4a78a575570d9500 Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 30 Dec 2024 13:26:44 -0600 Subject: [PATCH] add get_observation_categories --- .../api/routers/discovery_router.py | 10 +++- src/basic_memory/mcp/main.py | 4 +- src/basic_memory/mcp/tools/__init__.py | 2 + src/basic_memory/mcp/tools/discovery.py | 28 ++++++++- .../repository/observation_repository.py | 2 +- tests/api/test_discovery_router.py | 42 ++++++++++++- tests/mcp/test_tool_discovery.py | 59 ++++++++++++++++++- 7 files changed, 135 insertions(+), 12 deletions(-) diff --git a/src/basic_memory/api/routers/discovery_router.py b/src/basic_memory/api/routers/discovery_router.py index 2d277d5c..d7eab965 100644 --- a/src/basic_memory/api/routers/discovery_router.py +++ b/src/basic_memory/api/routers/discovery_router.py @@ -3,7 +3,7 @@ from fastapi import APIRouter from loguru import logger -from basic_memory.deps import EntityServiceDep +from basic_memory.deps import EntityServiceDep, ObservationServiceDep from basic_memory.schemas import EntityTypeList, ObservationCategoryList router = APIRouter(prefix="/discovery", tags=["discovery"]) @@ -15,3 +15,11 @@ async def get_entity_types(entity_service: EntityServiceDep) -> EntityTypeList: logger.debug("Getting all entity types") types = await entity_service.get_entity_types() return EntityTypeList(types=types) + + +@router.get("/observation-categories", response_model=ObservationCategoryList) +async def get_observation_categories(observation_service: ObservationServiceDep) -> ObservationCategoryList: + """Get list of all unique observation categories in the system.""" + logger.debug("Getting all observation categories") + categories = await observation_service.observation_categories() + return ObservationCategoryList(categories=categories) diff --git a/src/basic_memory/mcp/main.py b/src/basic_memory/mcp/main.py index 9c725b7e..80850cc2 100644 --- a/src/basic_memory/mcp/main.py +++ b/src/basic_memory/mcp/main.py @@ -11,8 +11,8 @@ from basic_memory.config import config from basic_memory.mcp.server import mcp # Import tools to register them -from basic_memory.mcp.tools import knowledge, search, documents -__all__ = ["mcp", "knowledge", "search", "documents"] +from basic_memory.mcp.tools import knowledge, search, documents, discovery +__all__ = ["mcp", "knowledge", "search", "documents", "discovery"] def setup_logging(home_dir: str = config.home, log_file: str = "basic-memory.log"): diff --git a/src/basic_memory/mcp/tools/__init__.py b/src/basic_memory/mcp/tools/__init__.py index 9c60671a..8c61a35a 100644 --- a/src/basic_memory/mcp/tools/__init__.py +++ b/src/basic_memory/mcp/tools/__init__.py @@ -37,6 +37,7 @@ from basic_memory.mcp.tools.documents import ( from basic_memory.mcp.tools.discovery import ( get_entity_types, + get_observation_categories, ) __all__ = [ @@ -62,4 +63,5 @@ __all__ = [ # Discovery tools "get_entity_types", + "get_observation_categories", ] \ No newline at end of file diff --git a/src/basic_memory/mcp/tools/discovery.py b/src/basic_memory/mcp/tools/discovery.py index 741cb419..293b4fd7 100644 --- a/src/basic_memory/mcp/tools/discovery.py +++ b/src/basic_memory/mcp/tools/discovery.py @@ -4,7 +4,7 @@ from typing import List from loguru import logger -from basic_memory.schemas import EntityTypeList +from basic_memory.schemas import EntityTypeList, ObservationCategoryList from basic_memory.mcp.async_client import client from basic_memory.mcp.server import mcp @@ -30,4 +30,28 @@ async def get_entity_types() -> List[str]: logger.debug("Getting all entity types") url = "/discovery/entity-types" response = await client.get(url) - return EntityTypeList.model_validate(response.json()).types + return EntityTypeList.model_validate(response.json()) + + +@mcp.tool() +async def get_observation_categories() -> List[str]: + """List all unique observation categories in use across the knowledge graph. + + Examples: + categories = await get_observation_categories() + + # Returns list of strings like: + # [ + # "tech", + # "design", + # "feature", + # "note" + # ] + + Returns: + List of unique observation category strings used in the knowledge graph + """ + logger.debug("Getting all observation categories") + url = "/discovery/observation-categories" + response = await client.get(url) + return ObservationCategoryList.model_validate(response.json()) diff --git a/src/basic_memory/repository/observation_repository.py b/src/basic_memory/repository/observation_repository.py index b147d79b..b58e4abb 100644 --- a/src/basic_memory/repository/observation_repository.py +++ b/src/basic_memory/repository/observation_repository.py @@ -36,5 +36,5 @@ class ObservationRepository(Repository[Observation]): async def observation_categories(self) -> Sequence[str]: """Return a list of all observation categories.""" query = select(Observation.category).distinct() - result = await self.execute_query(query) + result = await self.execute_query(query, use_query_options=False) return result.scalars().all() diff --git a/tests/api/test_discovery_router.py b/tests/api/test_discovery_router.py index c444bc51..66a264d6 100644 --- a/tests/api/test_discovery_router.py +++ b/tests/api/test_discovery_router.py @@ -4,9 +4,10 @@ import pytest import pytest_asyncio from httpx import AsyncClient -from basic_memory.models.knowledge import Entity +from basic_memory.models.knowledge import Entity, Observation from basic_memory.repository.entity_repository import EntityRepository -from basic_memory.schemas import EntityTypeList +from basic_memory.schemas import EntityTypeList, ObservationCategoryList + pytestmark = pytest.mark.asyncio @@ -21,6 +22,10 @@ async def test_entities(entity_repository: EntityRepository) -> list[Entity]: description="Core memory service", path_id="component/memory_service", file_path="component/memory_service.md", + observations=[ + Observation(category="tech", content="Using SQLite for storage"), + Observation(category="design", content="Local-first architecture"), + ] ), Entity( name="File Format", @@ -28,6 +33,10 @@ async def test_entities(entity_repository: EntityRepository) -> list[Entity]: description="File format spec", path_id="spec/file_format", file_path="spec/file_format.md", + observations=[ + Observation(category="feature", content="Support for frontmatter"), + Observation(category="tech", content="UTF-8 encoding"), + ] ), Entity( name="Technical Decision", @@ -35,6 +44,10 @@ async def test_entities(entity_repository: EntityRepository) -> list[Entity]: description="Architecture decision", path_id="decision/tech_choice", file_path="decision/tech_choice.md", + observations=[ + Observation(category="note", content="Team discussed options"), + Observation(category="design", content="Selected for scalability"), + ] ), ] @@ -42,7 +55,6 @@ async def test_entities(entity_repository: EntityRepository) -> list[Entity]: return created -@pytest.mark.asyncio async def test_get_entity_types(client: AsyncClient, test_entities): """Test getting list of entity types.""" # Get types @@ -64,3 +76,27 @@ async def test_get_entity_types(client: AsyncClient, test_entities): # Types should be unique assert len(data.types) == len(set(data.types)) + + +async def test_get_observation_categories(client: AsyncClient, test_entities): + """Test getting list of observation categories.""" + # Get categories + response = await client.get("/discovery/observation-categories") + assert response.status_code == 200 + + # Parse response + data = ObservationCategoryList.model_validate(response.json()) + + # Should have categories from test data + assert len(data.categories) > 0 + assert "tech" in data.categories + assert "design" in data.categories + assert "feature" in data.categories + assert "note" in data.categories + + # Categories should all be strings + assert isinstance(data.categories, list) + assert all(isinstance(c, str) for c in data.categories) + + # Categories should be unique + assert len(data.categories) == len(set(data.categories)) diff --git a/tests/mcp/test_tool_discovery.py b/tests/mcp/test_tool_discovery.py index 0459827e..3e9c8ce0 100644 --- a/tests/mcp/test_tool_discovery.py +++ b/tests/mcp/test_tool_discovery.py @@ -2,9 +2,10 @@ import pytest -from basic_memory.mcp.tools.discovery import get_entity_types -from basic_memory.schemas import Entity, CreateEntityRequest -from basic_memory.mcp.tools.knowledge import create_entities +from basic_memory.mcp.tools.discovery import get_entity_types, get_observation_categories +from basic_memory.schemas import Entity, CreateEntityRequest, AddObservationsRequest +from basic_memory.mcp.tools.knowledge import create_entities, add_observations +from basic_memory.schemas.request import ObservationCreate @pytest.mark.asyncio @@ -61,3 +62,55 @@ async def test_get_entity_types_empty(client): # Should return empty list, not error assert isinstance(result, list) assert len(result) == 0 + + +@pytest.mark.asyncio +async def test_get_observation_categories(client): + """Test getting list of observation categories.""" + # First create an entity with categorized observations + request = CreateEntityRequest( + entities=[ + Entity( + name="Test Entity", + entity_type="test", + path_id="test/entity", + description="Test entity", + observations=[] + ) + ] + ) + entity = (await create_entities(request)).entities[0] + + # Add observations with different categories + observations = [ + ObservationCreate(content="Technical detail", category="tech"), + ObservationCreate(content="Design decision", category="design"), + ObservationCreate(content="Feature spec", category="feature"), + ObservationCreate(content="General note", category="note") + ] + + await add_observations(AddObservationsRequest(path_id=entity.path_id, observations=observations)) + + # Get categories + result = await get_observation_categories() + + # Verify results + assert isinstance(result, list) + assert all(isinstance(c, str) for c in result) + assert "tech" in result + assert "design" in result + assert "feature" in result + assert "note" in result + + # Should be unique + assert len(result) == len(set(result)) + + +@pytest.mark.asyncio +async def test_get_observation_categories_empty(client): + """Test getting observation categories when no observations exist.""" + result = await get_observation_categories() + + # Should return empty list, not error + assert isinstance(result, list) + assert len(result) == 0