diff --git a/src/basic_memory/api/routers/discovery_router.py b/src/basic_memory/api/routers/discovery_router.py index d7eab965..772bca05 100644 --- a/src/basic_memory/api/routers/discovery_router.py +++ b/src/basic_memory/api/routers/discovery_router.py @@ -1,10 +1,12 @@ """Router for knowledge discovery and analytics operations.""" +from typing import Optional + from fastapi import APIRouter from loguru import logger from basic_memory.deps import EntityServiceDep, ObservationServiceDep -from basic_memory.schemas import EntityTypeList, ObservationCategoryList +from basic_memory.schemas import EntityTypeList, ObservationCategoryList, TypedEntityList, EntityResponse router = APIRouter(prefix="/discovery", tags=["discovery"]) @@ -23,3 +25,26 @@ async def get_observation_categories(observation_service: ObservationServiceDep) logger.debug("Getting all observation categories") categories = await observation_service.observation_categories() return ObservationCategoryList(categories=categories) + + +@router.get("/entities/{entity_type}", response_model=TypedEntityList) +async def list_entities_by_type( + entity_service: EntityServiceDep, + entity_type: str, + include_related: bool = False, + sort_by: Optional[str] = "updated_at", +) -> TypedEntityList: + """List all entities of a specific type.""" + logger.debug(f"Listing entities of type: {entity_type}") + entities = await entity_service.list_entities( + entity_type=entity_type, + sort_by=sort_by, + include_related=include_related + ) + return TypedEntityList( + entity_type=entity_type, + entities=[EntityResponse.model_validate(e) for e in entities], + total=len(entities), + sort_by=sort_by, + include_related=include_related + ) \ 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 293b4fd7..c72999a1 100644 --- a/src/basic_memory/mcp/tools/discovery.py +++ b/src/basic_memory/mcp/tools/discovery.py @@ -1,10 +1,10 @@ """Tools for discovering and analyzing knowledge graph structure.""" -from typing import List +from typing import List, Optional from loguru import logger -from basic_memory.schemas import EntityTypeList, ObservationCategoryList +from basic_memory.schemas import EntityTypeList, ObservationCategoryList, TypedEntityList from basic_memory.mcp.async_client import client from basic_memory.mcp.server import mcp @@ -55,3 +55,33 @@ async def get_observation_categories() -> List[str]: url = "/discovery/observation-categories" response = await client.get(url) return ObservationCategoryList.model_validate(response.json()) + + +@mcp.tool() +async def list_by_type( + entity_type: str, + include_related: bool = False, + sort_by: Optional[str] = "updated_at" +) -> TypedEntityList: + """List all entities of a specific type. + + Example: + # Get all features + features = await list_by_type("feature") + + # Get components with relations + components = await list_by_type( + "component", + include_related=True + ) + """ + logger.debug(f"Listing entities of type: {entity_type}") + params = { + "include_related": "true" if include_related else "false" + } + if sort_by: + params["sort_by"] = sort_by + + url = f"/discovery/entities/{entity_type}" + response = await client.get(url, params=params) + return TypedEntityList.model_validate(response.json()) \ No newline at end of file diff --git a/src/basic_memory/repository/entity_repository.py b/src/basic_memory/repository/entity_repository.py index e078d0ab..e7fc8b02 100644 --- a/src/basic_memory/repository/entity_repository.py +++ b/src/basic_memory/repository/entity_repository.py @@ -2,7 +2,7 @@ from typing import List, Optional, Sequence -from sqlalchemy import select, or_ +from sqlalchemy import select, or_, asc, desc from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from sqlalchemy.orm import selectinload from sqlalchemy.orm.interfaces import LoaderOption @@ -27,15 +27,22 @@ class EntityRepository(Repository[Entity]): self, entity_type: Optional[str] = None, doc_id: Optional[int] = None, + sort_by: Optional[str] = "updated_at", ) -> Sequence[Entity]: - """List all entities, optionally filtered by type.""" + """List all entities, optionally filtered by type and sorted.""" query = self.select().options(*self.get_load_options()) + # Apply filters if entity_type: query = query.where(Entity.entity_type == entity_type) if doc_id: query = query.where(Entity.doc_id == doc_id) + # Apply sorting + if sort_by: + sort_field = getattr(Entity, sort_by, Entity.updated_at) + query = query.order_by(asc(sort_field)) + result = await self.execute_query(query) return list(result.scalars().all()) diff --git a/src/basic_memory/schemas/__init__.py b/src/basic_memory/schemas/__init__.py index 6daf51be..f5ae7981 100644 --- a/src/basic_memory/schemas/__init__.py +++ b/src/basic_memory/schemas/__init__.py @@ -44,7 +44,7 @@ from basic_memory.schemas.response import ( # Discovery and analytics models from basic_memory.schemas.discovery import ( EntityTypeList, - ObservationCategoryList, + ObservationCategoryList, TypedEntityList, ) # For convenient imports, export all models @@ -76,4 +76,5 @@ __all__ = [ # Discovery and Analytics "EntityTypeList", "ObservationCategoryList", + "TypedEntityList" ] diff --git a/src/basic_memory/schemas/discovery.py b/src/basic_memory/schemas/discovery.py index 2977258d..280569c8 100644 --- a/src/basic_memory/schemas/discovery.py +++ b/src/basic_memory/schemas/discovery.py @@ -1,7 +1,9 @@ """Schemas for knowledge discovery and analytics endpoints.""" -from typing import List -from pydantic import BaseModel +from typing import List, Optional +from pydantic import BaseModel, Field + +from basic_memory.schemas.response import EntityResponse class EntityTypeList(BaseModel): @@ -12,3 +14,12 @@ class EntityTypeList(BaseModel): class ObservationCategoryList(BaseModel): """List of unique observation categories in the system.""" categories: List[str] + + +class TypedEntityList(BaseModel): + """List of entities of a specific type.""" + entity_type: str = Field(..., description="Type of entities in the list") + entities: List[EntityResponse] + total: int = Field(..., description="Total number of entities") + sort_by: Optional[str] = Field(None, description="Field used for sorting") + include_related: bool = Field(False, description="Whether related entities are included") \ No newline at end of file diff --git a/src/basic_memory/services/entity_service.py b/src/basic_memory/services/entity_service.py index 94cacd2a..783c821c 100644 --- a/src/basic_memory/services/entity_service.py +++ b/src/basic_memory/services/entity_service.py @@ -1,6 +1,6 @@ """Service for managing entities in the database.""" -from typing import Dict, Any, Sequence, List +from typing import Dict, Any, Sequence, List, Optional from loguru import logger @@ -72,6 +72,19 @@ class EntityService(BaseService[EntityRepository]): """Get list of all distinct entity types in the system.""" logger.debug("Getting all distinct entity types") return await self.repository.get_entity_types() + + async def list_entities( + self, + entity_type: Optional[str] = None, + sort_by: Optional[str] = "updated_at", + include_related: bool = False, + ) -> Sequence[EntityModel]: + """List entities with optional filtering and sorting.""" + logger.debug(f"Listing entities: type={entity_type} sort={sort_by}") + return await self.repository.list_entities( + entity_type=entity_type, + sort_by=sort_by + ) async def delete_entity(self, path_id: str) -> bool: """Delete entity from database.""" @@ -88,4 +101,4 @@ class EntityService(BaseService[EntityRepository]): """Delete entities and their files.""" logger.debug(f"Deleting entities: {path_ids}") deleted_count = await self.repository.delete_by_path_ids(path_ids) - return deleted_count > 0 + return deleted_count > 0 \ No newline at end of file diff --git a/tests/api/test_discovery_router.py b/tests/api/test_discovery_router.py index 66a264d6..300bd80a 100644 --- a/tests/api/test_discovery_router.py +++ b/tests/api/test_discovery_router.py @@ -6,7 +6,7 @@ from httpx import AsyncClient from basic_memory.models.knowledge import Entity, Observation from basic_memory.repository.entity_repository import EntityRepository -from basic_memory.schemas import EntityTypeList, ObservationCategoryList +from basic_memory.schemas import EntityTypeList, ObservationCategoryList, TypedEntityList pytestmark = pytest.mark.asyncio @@ -49,6 +49,17 @@ async def test_entities(entity_repository: EntityRepository) -> list[Entity]: Observation(category="design", content="Selected for scalability"), ] ), + # Add another technical component for sorting tests + Entity( + name="API Service", + entity_type="technical_component", + description="API layer", + path_id="component/api_service", + file_path="component/api_service.md", + observations=[ + Observation(category="tech", content="FastAPI based"), + ] + ), ] created = await entity_repository.add_all(entities) @@ -100,3 +111,51 @@ async def test_get_observation_categories(client: AsyncClient, test_entities): # Categories should be unique assert len(data.categories) == len(set(data.categories)) + + +async def test_list_entities_by_type(client: AsyncClient, test_entities): + """Test listing entities by type.""" + # List technical components + response = await client.get("/discovery/entities/technical_component") + assert response.status_code == 200 + + # Parse response + data = TypedEntityList.model_validate(response.json()) + + # Check response structure + assert data.entity_type == "technical_component" + assert len(data.entities) == 2 + assert data.total == 2 + + # Verify content + names = {e.name for e in data.entities} + assert "Memory Service" in names + assert "API Service" in names + + +async def test_list_entities_with_sorting(client: AsyncClient, test_entities): + """Test listing entities with different sort options.""" + # Sort by name + response = await client.get("/discovery/entities/technical_component?sort_by=name") + assert response.status_code == 200 + data = TypedEntityList.model_validate(response.json()) + names = [e.name for e in data.entities] + assert names == sorted(names) # Should be alphabetical + + # Sort by path_id + response = await client.get("/discovery/entities/technical_component?sort_by=path_id") + assert response.status_code == 200 + data = TypedEntityList.model_validate(response.json()) + path_ids = [e.path_id for e in data.entities] + assert path_ids == sorted(path_ids) + + +async def test_list_entities_empty_type(client: AsyncClient, test_entities): + """Test listing entities for a type that doesn't exist.""" + response = await client.get("/discovery/entities/nonexistent_type") + assert response.status_code == 200 + + data = TypedEntityList.model_validate(response.json()) + assert data.entity_type == "nonexistent_type" + assert len(data.entities) == 0 + assert data.total == 0 diff --git a/tests/mcp/test_tool_discovery.py b/tests/mcp/test_tool_discovery.py index a653b4e4..9fcdd0ba 100644 --- a/tests/mcp/test_tool_discovery.py +++ b/tests/mcp/test_tool_discovery.py @@ -2,16 +2,14 @@ import pytest -from basic_memory.mcp.tools.discovery import get_entity_types, get_observation_categories -from basic_memory.schemas import ( - Entity, - CreateEntityRequest, - AddObservationsRequest, - EntityTypeList, - ObservationCategoryList, +from basic_memory.mcp.tools.discovery import ( + get_entity_types, + get_observation_categories, + list_by_type, ) +from basic_memory.schemas import Entity, CreateEntityRequest, EntityTypeList, ObservationCategoryList from basic_memory.mcp.tools.knowledge import create_entities, add_observations -from basic_memory.schemas.request import ObservationCreate +from basic_memory.schemas.request import ObservationCreate, AddObservationsRequest @pytest.mark.asyncio @@ -25,22 +23,22 @@ async def test_get_entity_types(client): entity_type="technical_component", path_id="component/memory_service", description="Core memory service", - observations=["First observation"], + observations=["Using SQLite for storage", "Local-first architecture"] ), Entity( name="File Format", entity_type="specification", path_id="specification/file_format", description="File format spec", - observations=["Format details"], + observations=["Support for frontmatter", "UTF-8 encoding"] ), Entity( name="Tech Choice", entity_type="decision", path_id="decision/tech_choice", description="Technology decision", - observations=["Decision context"], - ), + observations=["Team discussed options", "Selected for scalability"] + ) ] ) await create_entities(request) @@ -74,8 +72,7 @@ async def test_get_observation_categories(client): entities=[ Entity( name="Test Entity", - entity_type="test", - path_id="test/entity", + entity_type="test_observation", description="Test entity", observations=[], ) @@ -107,6 +104,30 @@ async def test_get_observation_categories(client): +@pytest.mark.asyncio +async def test_list_by_type_with_sorting(client): + """Test listing entities with different sort options.""" + # Sort by name + result = await list_by_type("technical_component", sort_by="name") + names = [e.name for e in result.entities] + assert names == sorted(names) + + # Sort by path_id + result = await list_by_type("technical_component", sort_by="path_id") + path_ids = [e.path_id for e in result.entities] + assert path_ids == sorted(path_ids) + + +@pytest.mark.asyncio +async def test_list_by_type_empty(client): + """Test listing entities for a type that doesn't exist.""" + result = await list_by_type("nonexistent_type") + + assert result.entity_type == "nonexistent_type" + assert len(result.entities) == 0 + assert result.total == 0 + + @pytest.mark.asyncio async def test_get_observation_categories_empty(client): """Test getting observation categories when no observations exist."""