mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
add get_observation_categories
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user