diff --git a/src/basic_memory/api/routers/knowledge.py b/src/basic_memory/api/routers/knowledge.py index abac84a1..ca7bd0e5 100644 --- a/src/basic_memory/api/routers/knowledge.py +++ b/src/basic_memory/api/routers/knowledge.py @@ -9,14 +9,13 @@ from basic_memory.deps import ( ) from basic_memory.schemas import ( CreateEntityRequest, - CreateEntityResponse, + EntityListResponse, SearchNodesRequest, SearchNodesResponse, CreateRelationsRequest, EntityResponse, AddObservationsRequest, OpenNodesRequest, - OpenNodesResponse, DeleteEntitiesResponse, DeleteObservationsRequest, DeleteRelationsRequest, @@ -30,24 +29,24 @@ router = APIRouter(prefix="/knowledge", tags=["knowledge"]) ## Create endpoints -@router.post("/entities", response_model=CreateEntityResponse) +@router.post("/entities", response_model=EntityListResponse) async def create_entities( data: CreateEntityRequest, knowledge_service: KnowledgeServiceDep -) -> CreateEntityResponse: +) -> EntityListResponse: """Create new entities in the knowledge graph.""" entities = await knowledge_service.create_entities(data.entities) - return CreateEntityResponse( + return EntityListResponse( entities=[EntityResponse.model_validate(entity) for entity in entities] ) -@router.post("/relations", response_model=CreateEntityResponse) +@router.post("/relations", response_model=EntityListResponse) async def create_relations( data: CreateRelationsRequest, knowledge_service: KnowledgeServiceDep -) -> CreateEntityResponse: +) -> EntityListResponse: """Create relations between entities.""" updated_entities = await knowledge_service.create_relations(data.relations) - return CreateEntityResponse( + return EntityListResponse( entities=[EntityResponse.model_validate(entity) for entity in updated_entities] ) @@ -91,11 +90,11 @@ async def search_nodes( ) -@router.post("/nodes", response_model=OpenNodesResponse) -async def open_nodes(data: OpenNodesRequest, entity_service: EntityServiceDep) -> OpenNodesResponse: +@router.post("/nodes", response_model=EntityListResponse) +async def open_nodes(data: OpenNodesRequest, entity_service: EntityServiceDep) -> EntityListResponse: """Open specific nodes by their names.""" entities = await entity_service.open_nodes(data.entity_ids) - return OpenNodesResponse( + return EntityListResponse( entities=[EntityResponse.model_validate(entity) for entity in entities] ) @@ -122,12 +121,12 @@ async def delete_observations( return EntityResponse.model_validate(updated_entity) -@router.post("/relations/delete", response_model=CreateEntityResponse) +@router.post("/relations/delete", response_model=EntityListResponse) async def delete_relations( data: DeleteRelationsRequest, knowledge_service: KnowledgeServiceDep -) -> CreateEntityResponse: +) -> EntityListResponse: """Delete relations between entities.""" updated_entities = await knowledge_service.delete_relations(data.relations) - return CreateEntityResponse( + return EntityListResponse( entities=[EntityResponse.model_validate(entity) for entity in updated_entities] ) diff --git a/src/basic_memory/schemas/__init__.py b/src/basic_memory/schemas/__init__.py index 9399b4bd..38fbd867 100644 --- a/src/basic_memory/schemas/__init__.py +++ b/src/basic_memory/schemas/__init__.py @@ -36,9 +36,8 @@ from basic_memory.schemas.response import ( ObservationResponse, RelationResponse, EntityResponse, - CreateEntityResponse, + EntityListResponse, SearchNodesResponse, - OpenNodesResponse, DeleteEntitiesResponse, ) @@ -56,17 +55,13 @@ __all__ = [ "SearchNodesRequest", "OpenNodesRequest", "CreateRelationsRequest", - "DocumentCreateRequest", - "DocumentUpdateRequest", # Responses "SQLAlchemyModel", "ObservationResponse", - "ObservationsResponse", "RelationResponse", "EntityResponse", - "CreateEntityResponse", + "EntityListResponse", "SearchNodesResponse", - "OpenNodesResponse", "DeleteEntitiesResponse", # Delete Operations "DeleteEntitiesRequest", diff --git a/src/basic_memory/schemas/base.py b/src/basic_memory/schemas/base.py index a95b3fcb..781002a8 100644 --- a/src/basic_memory/schemas/base.py +++ b/src/basic_memory/schemas/base.py @@ -48,8 +48,10 @@ def to_snake_case(name: str) -> str: memory-service -> memory_service Memory_Service -> memory_service """ - # Replace spaces and hyphens with underscores - s1 = re.sub(r"[\s\-]", "_", name) + name = name.strip() + + # Replace spaces and hyphens and . with underscores + s1 = re.sub(r"[\s\-\\.]", "_", name) # Insert underscore between camelCase s2 = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", s1) @@ -58,7 +60,22 @@ def to_snake_case(name: str) -> str: return s2.lower() -PathId = Annotated[str, BeforeValidator(to_snake_case)] +def validate_path_format(path: str) -> str: + """Validate path has the correct format: type/name.""" + if not path or not isinstance(path, str): + raise ValueError("Path must be a non-empty string") + + parts = path.split('/') + if len(parts) != 2: + raise ValueError("Path must be in format: type/name") + + type_part, name_part = parts + if not type_part or not name_part: + raise ValueError("Both type and name must be non-empty") + + return path + +PathId = Annotated[str, BeforeValidator(to_snake_case), BeforeValidator(validate_path_format)] """Unique identifier in format '{path}/{normalized_name}'.""" Observation = Annotated[str, MinLen(1), MaxLen(1000)] diff --git a/src/basic_memory/schemas/response.py b/src/basic_memory/schemas/response.py index 5d523435..4b5f0e05 100644 --- a/src/basic_memory/schemas/response.py +++ b/src/basic_memory/schemas/response.py @@ -16,7 +16,7 @@ from typing import List, Optional, Dict, Any from pydantic import BaseModel, ConfigDict, Field, AliasPath, AliasChoices -from basic_memory.schemas.base import Observation, Relation, PathId +from basic_memory.schemas.base import Observation, Relation, PathId, Entity, EntityType class SQLAlchemyModel(BaseModel): @@ -38,7 +38,6 @@ class ObservationResponse(SQLAlchemyModel): Example Response: { - "id": 123, "content": "Implements SQLite storage for persistence" } """ @@ -112,19 +111,20 @@ class EntityResponse(SQLAlchemyModel): } """ + # Note this Class does not inherit form Entity because of the Entity.path_id semantics path_id: PathId name: str - entity_type: str + entity_type: EntityType description: Optional[str] = None observations: List[ObservationResponse] = [] relations: List[RelationResponse] = [] -class CreateEntityResponse(SQLAlchemyModel): +class EntityListResponse(SQLAlchemyModel): """Response for create_entities operation. - Returns complete information about all created entities, - including their generated IDs, initial observations, + Returns complete information about entities returned from the service, + including their path_ids, observations, and any established relations. Example Response: @@ -137,7 +137,6 @@ class CreateEntityResponse(SQLAlchemyModel): "description": "Knowledge graph search", "observations": [ { - "id": 125, "content": "Implements full-text search" } ], @@ -150,7 +149,6 @@ class CreateEntityResponse(SQLAlchemyModel): "description": "API Reference", "observations": [ { - "id": 126, "content": "Documents REST endpoints" } ], @@ -192,29 +190,6 @@ class SearchNodesResponse(SQLAlchemyModel): query: str -class OpenNodesResponse(SQLAlchemyModel): - """Response for retrieving specific entities. - - Returns complete Entity objects for all found entities. - Entities that don't exist are silently skipped. - - Example Response: - { - "entities": [ - { - "path_id": "component/memory_service", - "name": "MemoryService", - "entity_type": "component", - "description": "Core service", - "observations": [...], - "relations": [...] - } - ] - } - """ - - entities: List[EntityResponse] - class DeleteEntitiesResponse(SQLAlchemyModel): """Response indicating successful entity deletion. diff --git a/tests/api/test_knowledge_router.py b/tests/api/test_knowledge_router.py index 95f00095..f032b6a9 100644 --- a/tests/api/test_knowledge_router.py +++ b/tests/api/test_knowledge_router.py @@ -8,7 +8,7 @@ from httpx import AsyncClient from basic_memory.schemas import ( EntityResponse, - CreateEntityResponse, + EntityListResponse, ObservationResponse, RelationResponse, ) @@ -34,7 +34,7 @@ async def create_entity(client) -> EntityResponse: assert len(entity["observations"]) == 2 - create_response = CreateEntityResponse.model_validate(response_data) + create_response = EntityListResponse.model_validate(response_data) return create_response.entities[0] @@ -76,7 +76,7 @@ async def create_related_entities(client) -> List[RelationResponse]: # pyright: assert response.status_code == 200 data = response.json() - relation_response = CreateEntityResponse.model_validate(data) + relation_response = EntityListResponse.model_validate(data) assert len(relation_response.entities) == 2 source_entity = relation_response.entities[0] @@ -299,7 +299,7 @@ async def test_delete_relations(client, relation_repository): assert response.status_code == 200 data = response.json() - del_response = CreateEntityResponse.model_validate(data) + del_response = EntityListResponse.model_validate(data) assert len(del_response.entities) == 2 assert all(len(e.relations) == 0 for e in del_response.entities) @@ -346,7 +346,7 @@ async def test_delete_nonexistent_relations(client: AsyncClient): assert response.status_code == 200 data = response.json() - del_response = CreateEntityResponse.model_validate(data) + del_response = EntityListResponse.model_validate(data) assert del_response.entities == [] diff --git a/tests/schemas/test_schemas.py b/tests/schemas/test_schemas.py index df991dc9..b8607b5f 100644 --- a/tests/schemas/test_schemas.py +++ b/tests/schemas/test_schemas.py @@ -11,6 +11,7 @@ from basic_memory.schemas import ( SearchNodesRequest, OpenNodesRequest, RelationResponse, ) +from basic_memory.schemas.base import to_snake_case def test_entity_in_minimal(): @@ -53,10 +54,10 @@ def test_entity_in_validation(): def test_relation_in_validation(): """Test RelationIn validation.""" - data = {"from_id": "123", "to_id": "456", "relation_type": "test"} + data = {"from_id": "test/123", "to_id": "test/456", "relation_type": "test"} relation = Relation.model_validate(data) - assert relation.from_id == "123" - assert relation.to_id == "456" + assert relation.from_id == "test/123" + assert relation.to_id == "test/456" assert relation.relation_type == "test" assert relation.context is None @@ -71,10 +72,10 @@ def test_relation_in_validation(): def test_relation_response(): """Test RelationResponse validation.""" - data = {"from_id": 123, "to_id": 456, "relation_type": "test", "from_entity":{"path_id": "123"}, "to_entity":{"path_id": "456"}} + data = {"from_id": "test/123", "to_id": "test/456", "relation_type": "test", "from_entity":{"path_id": "test/123"}, "to_entity":{"path_id": "test/456"}} relation = RelationResponse.model_validate(data) - assert relation.from_id == "123" - assert relation.to_id == "456" + assert relation.from_id == "test/123" + assert relation.to_id == "test/456" assert relation.relation_type == "test" assert relation.context is None @@ -154,9 +155,86 @@ def test_search_nodes_input(): def test_open_nodes_input(): """Test OpenNodesInput validation.""" - open_input = OpenNodesRequest.model_validate({"entity_ids": ["test", "test2"]}) + open_input = OpenNodesRequest.model_validate({"entity_ids": ["test/test", "test/test2"]}) assert len(open_input.entity_ids) == 2 # Empty names list should fail with pytest.raises(ValidationError): OpenNodesRequest.model_validate({"entity_ids": []}) + + +def test_path_sanitization(): + """Test to_snake_case() handles various inputs correctly.""" + test_cases = [ + ("BasicMemory", "basic_memory"), # CamelCase + ("Memory Service", "memory_service"), # Spaces + ("memory-service", "memory_service"), # Hyphens + ("Memory_Service", "memory_service"), # Already has underscore + ("API2Service", "api2_service"), # Numbers + (" Spaces ", "spaces"), # Extra spaces + ("mixedCase", "mixed_case"), # Mixed case + ("snake_case_already", "snake_case_already"), # Already snake case + ("ALLCAPS", "allcaps"), # All caps + ("with.dots", "with_dots"), # Dots + ] + + for input_str, expected in test_cases: + result = to_snake_case(input_str) + assert result == expected, f"Failed for input: {input_str}" + + +def test_path_id_generation(): + """Test path_id property generates correct paths.""" + test_cases = [ + ( + {"name": "BasicMemory", "entity_type": "Project"}, + "project/basic_memory" + ), + ( + {"name": "Memory Service", "entity_type": "Component"}, + "component/memory_service" + ), + ( + {"name": "API Gateway", "entity_type": "Service"}, + "service/api_gateway" + ), + ( + {"name": "TestCase1", "entity_type": "Test"}, + "test/test_case1" + ), + ] + + for input_data, expected_path in test_cases: + entity = Entity.model_validate(input_data) + assert entity.path_id == expected_path, f"Failed for input: {input_data}" + + +def test_path_id_validation(): + """Test path ID format validation.""" + valid_paths = [ + "project/basic_memory", + "test/test_case_1", + "component/api_gateway", + ] + + invalid_paths = [ + "no_separator", # Missing / + "/missing_type", # Missing type + "type/", # Missing name + "type//double", # Double separator + "../path/traversal", # Path traversal attempt + "type/name/extra", # Too many parts + "", # Empty string + ] + + # Test valid paths + for path in valid_paths: + try: + Relation.model_validate({"from_id": path, "to_id": path, "relation_type": "test"}) + except ValidationError as e: + assert False, f"Valid path {path} failed validation: {e}" + + # Test invalid paths + for path in invalid_paths: + with pytest.raises(ValidationError): + Relation.model_validate({"from_id": path, "to_id": "test/valid", "relation_type": "test"}) \ No newline at end of file