refactor knowledge_service.create_relations

This commit is contained in:
phernandez
2024-12-23 10:01:04 -06:00
parent 76e3466bcf
commit a5bb6f8b0c
3 changed files with 30 additions and 33 deletions
-8
View File
@@ -40,11 +40,7 @@ from basic_memory.schemas.response import (
CreateEntityResponse,
SearchNodesResponse,
OpenNodesResponse,
AddObservationsResponse,
CreateRelationsResponse,
DeleteEntitiesResponse,
DeleteRelationsResponse,
DeleteObservationsResponse,
)
# For convenient imports, export all models
@@ -70,11 +66,7 @@ __all__ = [
"CreateEntityResponse",
"SearchNodesResponse",
"OpenNodesResponse",
"AddObservationsResponse",
"CreateRelationsResponse",
"DeleteEntitiesResponse",
"DeleteRelationsResponse",
"DeleteObservationsResponse",
# Delete Operations
"DeleteEntitiesRequest",
"DeleteRelationsRequest",
@@ -4,6 +4,7 @@ from typing import Sequence, List
from loguru import logger
from basic_memory.models import Entity as EntityModel
from basic_memory.schemas import Relation as RelationSchema
from basic_memory.services.exceptions import EntityNotFoundError
from basic_memory.services.relation_service import RelationService
@@ -17,31 +18,41 @@ class RelationOperations(EntityOperations):
super().__init__(*args, **kwargs)
self.relation_service = relation_service
async def create_relations(self, relations: List[RelationSchema]) -> Sequence[RelationSchema]:
"""Create relations and update affected entity files."""
async def create_relations(self, relations: List[RelationSchema]) -> Sequence[EntityModel]:
"""Create relations and return updated entities."""
logger.debug(f"Creating {len(relations)} relations")
created = []
updated_entities = []
update_entity_ids = set()
for relation in relations:
try:
# Create relation in DB
db_relation = await self.relation_service.create_relation(relation)
await self.relation_service.create_relation(relation)
# Update files with their new relations
for entity_id in [relation.from_id, relation.to_id]:
# 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)
await self.entity_service.update_entity(entity_id, {"checksum": checksum})
created.append(db_relation)
# Keep track of entities we need to update
update_entity_ids.add(relation.from_id)
update_entity_ids.add(relation.to_id)
except Exception as e:
logger.error(f"Failed to create relation: {e}")
continue
return created
# Get fresh copies of all updated entities
for entity_id in update_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