From 73c0a4078e539ef0e33e41ade35f65f78866d352 Mon Sep 17 00:00:00 2001 From: phernandez Date: Sun, 22 Dec 2024 12:52:01 -0600 Subject: [PATCH] more test fixing for entity int id --- src/basic_memory/api/routers/knowledge.py | 16 ++++++++-------- src/basic_memory/mcp/server.py | 4 ++-- src/basic_memory/schemas/__init__.py | 8 ++++---- src/basic_memory/schemas/delete.py | 2 +- src/basic_memory/schemas/response.py | 4 ++-- tests/api/test_knowledge_router.py | 8 ++++---- tests/schemas/test_schemas.py | 8 ++++---- 7 files changed, 25 insertions(+), 25 deletions(-) diff --git a/src/basic_memory/api/routers/knowledge.py b/src/basic_memory/api/routers/knowledge.py index cc3c3068..6b0ccade 100644 --- a/src/basic_memory/api/routers/knowledge.py +++ b/src/basic_memory/api/routers/knowledge.py @@ -16,14 +16,14 @@ from basic_memory.schemas import ( ObservationResponse, OpenNodesRequest, OpenNodesResponse, - DeleteEntityResponse, + DeleteEntitiesResponse, DeleteObservationsRequest, DeleteObservationsResponse, DeleteRelationsRequest, DeleteRelationsResponse, AddObservationsResponse, RelationResponse, - DeleteEntityRequest, + DeleteEntitiesRequest, ) from basic_memory.services.exceptions import EntityNotFoundError @@ -74,8 +74,8 @@ async def add_observations( ## Read endpoints -@router.get("/entities/{entity_id:path}", response_model=EntityResponse) -async def get_entity(entity_id: str, entity_service: EntityServiceDep) -> EntityResponse: +@router.get("/entities/{entity_id}", response_model=EntityResponse) +async def get_entity(entity_id: int, entity_service: EntityServiceDep) -> EntityResponse: """Get a specific entity by ID.""" try: entity = await entity_service.get_entity(entity_id) @@ -110,13 +110,13 @@ async def open_nodes(data: OpenNodesRequest, entity_service: EntityServiceDep) - ## Delete endpoints -@router.post("/entities/delete", response_model=DeleteEntityResponse) +@router.post("/entities/delete", response_model=DeleteEntitiesResponse) async def delete_entity( - data: DeleteEntityRequest, entity_service: EntityServiceDep -) -> DeleteEntityResponse: + data: DeleteEntitiesRequest, entity_service: EntityServiceDep +) -> DeleteEntitiesResponse: """Delete a specific entity by ID.""" deleted = await entity_service.delete_entities(data.entity_ids) - return DeleteEntityResponse(deleted=deleted) + return DeleteEntitiesResponse(deleted=deleted) @router.post("/observations/delete", response_model=DeleteObservationsResponse) diff --git a/src/basic_memory/mcp/server.py b/src/basic_memory/mcp/server.py index 86559646..a6b250b7 100644 --- a/src/basic_memory/mcp/server.py +++ b/src/basic_memory/mcp/server.py @@ -33,7 +33,7 @@ from basic_memory.schemas import ( OpenNodesRequest, AddObservationsRequest, CreateRelationsRequest, - DeleteEntityRequest, + DeleteEntitiesRequest, DeleteObservationsRequest, DeleteRelationsRequest, ) @@ -82,7 +82,7 @@ async def handle_list_tools() -> List[Tool]: Tool( name="delete_entities", description="Delete entities", - inputSchema=DeleteEntityRequest.model_json_schema(), + inputSchema=DeleteEntitiesRequest.model_json_schema(), ), Tool( name="delete_observations", diff --git a/src/basic_memory/schemas/__init__.py b/src/basic_memory/schemas/__init__.py index 9e274752..ba67f396 100644 --- a/src/basic_memory/schemas/__init__.py +++ b/src/basic_memory/schemas/__init__.py @@ -16,7 +16,7 @@ from basic_memory.schemas.base import ( # Delete operation models from basic_memory.schemas.delete import ( - DeleteEntityRequest, + DeleteEntitiesRequest, DeleteRelationsRequest, DeleteObservationsRequest, ) @@ -42,7 +42,7 @@ from basic_memory.schemas.response import ( OpenNodesResponse, AddObservationsResponse, CreateRelationsResponse, - DeleteEntityResponse, + DeleteEntitiesResponse, DeleteRelationsResponse, DeleteObservationsResponse, ) @@ -72,11 +72,11 @@ __all__ = [ "OpenNodesResponse", "AddObservationsResponse", "CreateRelationsResponse", - "DeleteEntityResponse", + "DeleteEntitiesResponse", "DeleteRelationsResponse", "DeleteObservationsResponse", # Delete Operations - "DeleteEntityRequest", + "DeleteEntitiesRequest", "DeleteRelationsRequest", "DeleteObservationsRequest", ] diff --git a/src/basic_memory/schemas/delete.py b/src/basic_memory/schemas/delete.py index e95f6cd1..5d2bf97a 100644 --- a/src/basic_memory/schemas/delete.py +++ b/src/basic_memory/schemas/delete.py @@ -24,7 +24,7 @@ from pydantic import BaseModel from basic_memory.schemas.base import Relation, Observation -class DeleteEntityRequest(BaseModel): +class DeleteEntitiesRequest(BaseModel): """Delete one or more entities from the knowledge graph. This operation: diff --git a/src/basic_memory/schemas/response.py b/src/basic_memory/schemas/response.py index affd0548..f4c02b8b 100644 --- a/src/basic_memory/schemas/response.py +++ b/src/basic_memory/schemas/response.py @@ -129,7 +129,7 @@ class EntityResponse(SQLAlchemyModel): } """ - id: str + id: int name: str entity_type: str description: Optional[str] = None @@ -282,7 +282,7 @@ class CreateRelationsResponse(SQLAlchemyModel): relations: List[Relation] -class DeleteEntityResponse(SQLAlchemyModel): +class DeleteEntitiesResponse(SQLAlchemyModel): """Response indicating successful entity deletion. A simple boolean response confirming the delete operation diff --git a/tests/api/test_knowledge_router.py b/tests/api/test_knowledge_router.py index bb8f543c..3ed15a3f 100644 --- a/tests/api/test_knowledge_router.py +++ b/tests/api/test_knowledge_router.py @@ -5,7 +5,6 @@ from typing import List import pytest from httpx import AsyncClient -from basic_memory.models import Entity from basic_memory.schemas import ( EntityResponse, CreateEntityResponse, @@ -30,7 +29,6 @@ async def create_entity(client) -> EntityResponse: assert len(response_data["entities"]) == 1 entity = response_data["entities"][0] - assert entity["id"] == Entity.generate_id(entity["entity_type"], entity["name"]) assert entity["name"] == data["name"] entity_type = entity.get("entity_type") assert entity_type == data["entity_type"] @@ -182,14 +180,16 @@ async def test_open_nodes(client: AsyncClient): {"name": "Alpha Test", "entity_type": "test"}, {"name": "Beta Test", "entity_type": "test"}, ] - await client.post("/knowledge/entities", json={"entities": entities}) + create_response = await client.post("/knowledge/entities", json={"entities": entities}) + created_entities = create_response.json() + assert len(created_entities["entities"]) == 2 # open nodes response = await client.post( "/knowledge/nodes", json={ "entity_ids": [ - "test/alpha_test", + created_entities["entities"][0]["id"], ] }, ) diff --git a/tests/schemas/test_schemas.py b/tests/schemas/test_schemas.py index e3a3f81c..d16b140c 100644 --- a/tests/schemas/test_schemas.py +++ b/tests/schemas/test_schemas.py @@ -53,10 +53,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": 123, "to_id": 456, "relation_type": "test"} relation = Relation.model_validate(data) - assert relation.from_id == "123" - assert relation.to_id == "456" + assert relation.from_id == 123 + assert relation.to_id == 456 assert relation.relation_type == "test" assert relation.context is None @@ -146,7 +146,7 @@ def test_search_nodes_input(): def test_open_nodes_input(): """Test OpenNodesInput validation.""" - open_input = OpenNodesRequest.model_validate({"entity_ids": ["entity1", "entity2"]}) + open_input = OpenNodesRequest.model_validate({"entity_ids": [1, 2]}) assert len(open_input.entity_ids) == 2 # Empty names list should fail