refactor knowledge_service to return Entity objects on updates

This commit is contained in:
phernandez
2024-12-23 10:11:56 -06:00
parent a5bb6f8b0c
commit cbbfa6e637
4 changed files with 89 additions and 26 deletions
+6 -9
View File
@@ -45,8 +45,6 @@ async def create_relations(
data: CreateRelationsRequest, knowledge_service: KnowledgeServiceDep
) -> CreateEntityResponse:
"""Create relations between entities."""
# TODO knowledge_service.create_relations should return updated Entities
updated_entities = await knowledge_service.create_relations(data.relations)
return CreateEntityResponse(
entities=[EntityResponse.model_validate(entity) for entity in updated_entities]
@@ -119,16 +117,14 @@ async def delete_observations(
) -> EntityResponse:
"""Delete observations from an entity."""
entity_id = data.entity_id
# TODO add knowledge_service.delete_observations
updated_entity = await knowledge_service.delete_observations(entity_id, data.deletions)
return EntityResponse.model_validate(updated_entity)
@router.post("/relations/delete", response_model=EntityResponse)
@router.post("/relations/delete", response_model=CreateEntityResponse)
async def delete_relations(
data: DeleteRelationsRequest, knowledge_service: KnowledgeServiceDep
) -> EntityResponse:
) -> CreateEntityResponse:
"""Delete relations between entities."""
to_delete = [
{
@@ -138,6 +134,7 @@ async def delete_relations(
}
for relation in data.relations
]
# TODO add knowledge_service.delete_relations
updated_entity = await knowledge_service.delete_relations(to_delete)
return EntityResponse.model_validate(updated_entity)
updated_entities = await knowledge_service.delete_relations(to_delete)
return CreateEntityResponse(
entities=[EntityResponse.model_validate(entity) for entity in updated_entities]
)
@@ -24,7 +24,7 @@ class ObservationOperations(RelationOperations):
logger.debug(f"Adding observations to entity {entity_id}")
try:
# Get updated entity
# Get entity to update
entity = await self.entity_service.get_entity(entity_id)
if not entity:
raise EntityNotFoundError(f"Entity not found: {entity_id}")
@@ -32,15 +32,39 @@ class ObservationOperations(RelationOperations):
# Add observations to DB
await self.observation_service.add_observations(entity_id, observations, context)
# Get updated entity
updated_entity = await self.entity_service.get_entity(entity_id)
# Write updated file and checksum
checksum = await self.write_entity_file(entity)
await self.entity_service.update_entity(entity_id, {"checksum": checksum})
return updated_entity
except Exception as e:
logger.error(f"Failed to add observations: {e}")
raise
async def delete_observations(self, entity_id: int, observations: List[str]) -> EntityModel:
"""Delete observations from entity and update its file."""
logger.debug(f"Deleting observations from entity {entity_id}")
try:
# Get updated entity
entity = await self.entity_service.get_entity(entity_id)
if not entity:
raise EntityNotFoundError(f"Entity not found: {entity_id}")
# Delete observations from DB
await self.observation_service.delete_observations(entity_id, observations)
# Write updated file
checksum = await self.write_entity_file(entity)
# Update checksum in DB
await self.entity_service.update_entity(entity_id, {"checksum": checksum})
# Get final entity with all updates
return await self.entity_service.get_entity(entity_id)
except Exception as e:
logger.error(f"Failed to add observations: {e}")
logger.error(f"Failed to delete observations: {e}")
raise
@@ -1,6 +1,6 @@
"""Relation operations for knowledge service."""
from typing import Sequence, List
from typing import Sequence, List, Dict, Any
from loguru import logger
@@ -21,8 +21,8 @@ class RelationOperations(EntityOperations):
async def create_relations(self, relations: List[RelationSchema]) -> Sequence[EntityModel]:
"""Create relations and return updated entities."""
logger.debug(f"Creating {len(relations)} relations")
updated_entities = []
update_entity_ids = set()
created_entities = []
updated_entity_ids = set()
for relation in relations:
try:
@@ -30,15 +30,15 @@ class RelationOperations(EntityOperations):
await self.relation_service.create_relation(relation)
# Keep track of entities we need to update
update_entity_ids.add(relation.from_id)
update_entity_ids.add(relation.to_id)
updated_entity_ids.add(relation.from_id)
updated_entity_ids.add(relation.to_id)
except Exception as e:
logger.error(f"Failed to create relation: {e}")
continue
# Get fresh copies of all updated entities
for entity_id in update_entity_ids:
for entity_id in updated_entity_ids:
try:
# Get fresh entity
entity = await self.entity_service.get_entity(entity_id)
@@ -49,10 +49,50 @@ class RelationOperations(EntityOperations):
checksum = await self.write_entity_file(entity)
updated = await self.entity_service.update_entity(entity_id, {"checksum": checksum})
updated_entities.append(updated)
created_entities.append(updated)
except Exception as e:
logger.error(f"Failed to update entity {entity_id}: {e}")
continue
return updated_entities
return created_entities
async def delete_relations(self, to_delete: List[Dict[str, Any]]) -> Sequence[EntityModel]:
"""Delete relations and return all updated entities."""
logger.debug(f"Deleting {len(to_delete)} relations")
updated_entity_ids = set()
try:
# Delete relations from DB
for relation in to_delete:
updated_entity_ids.add(relation["from_id"])
updated_entity_ids.add(relation["to_id"])
deleted = await self.relation_service.delete_relations(to_delete)
if not deleted:
logger.warning("No relations were deleted")
# Get fresh copies of all updated entities
updated_entities = []
for entity_id in updated_entity_ids:
try:
# Get fresh entity
entity = await self.entity_service.get_entity(entity_id)
if not entity:
raise EntityNotFoundError(f"Entity not found: {entity_id}")
# Write updated file
checksum = await self.write_entity_file(entity)
updated = await self.entity_service.update_entity(entity_id, {"checksum": checksum})
updated_entities.append(updated)
except Exception as e:
logger.error(f"Failed to update entity {entity_id}: {e}")
continue
return updated_entities
except Exception as e:
logger.error(f"Failed to delete relations: {e}")
raise
+7 -5
View File
@@ -99,15 +99,17 @@ async def test_add_observations(knowledge_service: KnowledgeService):
# Add observations
observations = ["Test observation 1", "Test observation 2"]
updated = await knowledge_service.add_observations(entity.id, observations, "Test context")
updated_entity = await knowledge_service.add_observations(
entity.id, observations, "Test context"
)
# Verify observations in DB
assert len(updated.observations) == 2
assert updated.observations[0].content == "Test observation 1"
assert updated.observations[1].content == "Test observation 2"
assert len(updated_entity.observations) == 2
assert updated_entity.observations[0].content == "Test observation 1"
assert updated_entity.observations[1].content == "Test observation 2"
# Verify file was updated
file_path = knowledge_service.get_entity_path(updated)
file_path = knowledge_service.get_entity_path(updated_entity)
content, _ = await knowledge_service.file_service.read_file(file_path)
for obs in observations:
assert obs in content