add get_observation_categories

This commit is contained in:
phernandez
2024-12-30 13:26:44 -06:00
parent 1dfcc0bc4b
commit 2bc35f34e3
7 changed files with 135 additions and 12 deletions
@@ -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)
+2 -2
View File
@@ -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"):
+2
View File
@@ -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",
]
+26 -2
View File
@@ -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())
@@ -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()
+39 -3
View File
@@ -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))
+56 -3
View File
@@ -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