mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
refactor knowledge_service to return Entity objects on updates
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user