fix schema plural namings

This commit is contained in:
phernandez
2024-12-14 19:30:01 -06:00
parent 635c2851a7
commit b471e1f8d7
6 changed files with 44 additions and 44 deletions
+11 -11
View File
@@ -4,26 +4,26 @@ from fastapi import APIRouter
from basic_memory.deps import MemoryServiceDep
from basic_memory.schemas import (
CreateEntitiesRequest, CreateEntitiesResponse,
CreateEntityRequest, CreateEntityResponse,
SearchNodesRequest, SearchNodesResponse,
CreateRelationsRequest, CreateRelationsResponse,
EntityResponse, RelationResponse, AddObservationsRequest, ObservationResponse,
OpenNodesRequest, OpenNodesResponse,
DeleteEntitiesResponse,
DeleteEntityResponse,
DeleteObservationsRequest, DeleteObservationsResponse, AddObservationsResponse
)
router = APIRouter(prefix="/knowledge", tags=["knowledge"])
@router.post("/entities", response_model=CreateEntitiesResponse)
@router.post("/entities", response_model=CreateEntityResponse)
async def create_entities(
data: CreateEntitiesRequest,
data: CreateEntityRequest,
memory_service: MemoryServiceDep
) -> CreateEntitiesResponse:
) -> CreateEntityResponse:
"""Create new entities in the knowledge graph."""
entities = await memory_service.create_entities(data.entities)
return CreateEntitiesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities])
return CreateEntityResponse(entities=[EntityResponse.model_validate(entity) for entity in entities])
@router.get("/entities/{entity_id:path}", response_model=EntityResponse)
@@ -36,14 +36,14 @@ async def get_entity(
return EntityResponse.model_validate(entity)
@router.delete("/entities/{entity_id}", response_model=DeleteEntitiesResponse)
@router.delete("/entities/{entity_id}", response_model=DeleteEntityResponse)
async def delete_entity(
entity_id: str,
memory_service: MemoryServiceDep
) -> DeleteEntitiesResponse:
) -> DeleteEntityResponse:
"""Delete a specific entity by ID."""
deleted = await memory_service.delete_entities([entity_id])
return DeleteEntitiesResponse(deleted=deleted) # pyright: ignore [reportArgumentType]
return DeleteEntityResponse(deleted=deleted) # pyright: ignore [reportArgumentType]
@router.post("/nodes", response_model=OpenNodesResponse)
@@ -66,11 +66,11 @@ async def create_relations(
return CreateRelationsResponse(relations=[RelationResponse.model_validate(relation) for relation in relations])
@router.delete("/relations/{relation_id}", response_model=DeleteEntitiesResponse)
@router.delete("/relations/{relation_id}", response_model=DeleteEntityResponse)
async def delete_relation(
relation_id: int,
memory_service: MemoryServiceDep
) -> DeleteEntitiesResponse:
) -> DeleteEntityResponse:
"""Delete a specific relation by ID."""
# TODO: Implement delete_relation in memory service
raise NotImplementedError("Delete relation not implemented yet")
+10 -10
View File
@@ -21,12 +21,12 @@ from basic_memory.repository.observation_repository import ObservationRepository
from basic_memory.repository.relation_repository import RelationRepository
from basic_memory.schemas import (
# Tool inputs
CreateEntitiesRequest, SearchNodesRequest, OpenNodesRequest,
CreateRelationsRequest, DeleteEntitiesRequest,
CreateEntityRequest, SearchNodesRequest, OpenNodesRequest,
CreateRelationsRequest, DeleteEntityRequest,
DeleteObservationsRequest,
# Tool responses
CreateEntitiesResponse, SearchNodesResponse, OpenNodesResponse,
AddObservationsResponse, CreateRelationsResponse, DeleteEntitiesResponse,
CreateEntityResponse, SearchNodesResponse, OpenNodesResponse,
AddObservationsResponse, CreateRelationsResponse, DeleteEntityResponse,
EntityResponse, ObservationResponse, RelationResponse, AddObservationsRequest
)
from basic_memory.services import EntityService, ObservationService, RelationService
@@ -103,7 +103,7 @@ async def handle_create_entities(
"""Handle create_entities tool call."""
# Validate input
logger.debug(f"Creating entities with args: {args}")
input_args = CreateEntitiesRequest.model_validate(args)
input_args = CreateEntityRequest.model_validate(args)
logger.debug(f"Validated input: {len(input_args.entities)} entities")
# Call service with validated data
@@ -111,7 +111,7 @@ async def handle_create_entities(
logger.debug(f"Created {len(entities)} entities")
# Format response
response = CreateEntitiesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities])
response = CreateEntityResponse(entities=[EntityResponse.model_validate(entity) for entity in entities])
logger.debug("Formatted create_entities response")
return create_response(response)
@@ -192,10 +192,10 @@ async def handle_delete_entities(
) -> EmbeddedResource:
"""Handle delete_entities tool call."""
logger.debug(f"Deleting entities: {args}")
input_args = DeleteEntitiesRequest.model_validate(args)
input_args = DeleteEntityRequest.model_validate(args)
deleted = await service.delete_entities(input_args.names)
logger.debug(f"Deleted entities: {deleted}")
response = DeleteEntitiesResponse(deleted=deleted)
response = DeleteEntityResponse(deleted=deleted)
return create_response(response)
@@ -240,7 +240,7 @@ class MemoryServer(Server):
Tool(
name="create_entities",
description="Create multiple new entities in the knowledge graph",
inputSchema=CreateEntitiesRequest.model_json_schema()
inputSchema=CreateEntityRequest.model_json_schema()
),
Tool(
name="search_nodes",
@@ -265,7 +265,7 @@ class MemoryServer(Server):
Tool(
name="delete_entities",
description="Delete entities from the knowledge graph",
inputSchema=DeleteEntitiesRequest.model_json_schema()
inputSchema=DeleteEntityRequest.model_json_schema()
),
Tool(
name="delete_observations",
+11 -11
View File
@@ -76,33 +76,33 @@ class EntityResponse(EntityBase, SQLAlchemyModel):
relations: List[RelationResponse] = []
model_config = ConfigDict(populate_by_name=True)
# Tool Input Schemas
class CreateEntitiesRequest(BaseModel):
"""Input schema for create_entities tool."""
# Tool Request schemas
class CreateEntityRequest(BaseModel):
"""Request schema for create_entities tool."""
entities: Annotated[List[EntityRequest], Len(min_length=1)]
class SearchNodesRequest(BaseModel):
"""Input schema for search_nodes tool."""
"""Request schema for search_nodes tool."""
query: str
class OpenNodesRequest(BaseModel):
"""Input schema for open_nodes tool."""
"""Request schema for open_nodes tool."""
names: Annotated[List[str], Len(min_length=1)]
class CreateRelationsRequest(BaseModel):
"""Input schema for create_relations tool."""
"""Request schema for create_relations tool."""
relations: List[RelationRequest]
class DeleteEntitiesRequest(BaseModel):
"""Input schema for delete_entities tool."""
class DeleteEntityRequest(BaseModel):
"""Request schema for delete_entities tool."""
names: List[str]
class DeleteObservationsRequest(BaseModel):
"""Input schema for delete_observations tool."""
"""Request schema for delete_observations tool."""
entity_id: str
deletions: List[str] # TODO: Make this more specific
class CreateEntitiesResponse(SQLAlchemyModel):
class CreateEntityResponse(SQLAlchemyModel):
"""Response for create_entities tool."""
entities: List[EntityResponse]
@@ -124,7 +124,7 @@ class CreateRelationsResponse(SQLAlchemyModel):
"""Response for create_relations tool."""
relations: List[RelationResponse]
class DeleteEntitiesResponse(SQLAlchemyModel):
class DeleteEntityResponse(SQLAlchemyModel):
"""Response for delete_entities tool."""
deleted: List[str]