refactor schemas and services

This commit is contained in:
phernandez
2024-12-07 22:15:22 -06:00
parent ac95036d0e
commit c71dd2cf0d
8 changed files with 159 additions and 178 deletions
+5 -5
View File
@@ -3,8 +3,8 @@ from datetime import datetime, UTC
from pathlib import Path
from basic_memory.repository import EntityRepository
from basic_memory.schemas import Entity
from basic_memory.models import Entity as EntityModel
from basic_memory.schemas import EntityIn
from basic_memory.models import Entity
from . import ServiceError
class EntityService:
@@ -17,7 +17,7 @@ class EntityService:
self.project_path = project_path
self.entity_repo = entity_repo
async def create_entity(self, entity: Entity) -> EntityModel:
async def create_entity(self, entity: EntityIn) -> Entity:
"""Create a new entity in the database."""
# Create DB record
db_data = {
@@ -26,7 +26,7 @@ class EntityService:
}
return await self.entity_repo.create(db_data)
async def get_entity(self, entity_id: str) -> EntityModel:
async def get_entity(self, entity_id: str) -> Entity:
"""Get entity by ID."""
db_entity = await self.entity_repo.find_by_id(entity_id)
if not db_entity:
@@ -35,7 +35,7 @@ class EntityService:
return db_entity
# TODO name is not uniaue
async def get_by_name(self, name: str) -> EntityModel:
async def get_by_name(self, name: str) -> Entity:
"""Get entity by name."""
db_entity = await self.entity_repo.find_by_name(name)
if not db_entity:
+37 -32
View File
@@ -3,9 +3,9 @@ import asyncio
from typing import List, Dict, Any, Optional
from pathlib import Path
from basic_memory.models import Entity, Observation
from basic_memory.schemas import (
Entity, Observation, Relation,
ObservationsIn, ObservationsOut, ObservationOut
ObservationsIn, ObservationsOut, ObservationOut, EntityIn, RelationIn, RelationOut
)
from basic_memory.fileio import write_entity_file, read_entity_file, delete_entity_file
from basic_memory.services import EntityService, RelationService, ObservationService
@@ -28,24 +28,36 @@ class MemoryService:
async def create_entities(self, entities_data: List[Dict[str, Any]]) -> List[Entity]:
"""Create multiple entities with their observations."""
entities = [Entity.model_validate(data) for data in entities_data]
entities_in = [EntityIn.model_validate(data) for data in entities_data]
print(f"\nCreating entities with observations:")
for e in entities_in:
print(f"Entity {e.name}: {len(e.observations)} observations")
# Write files in parallel (filesystem is source of truth)
async def write_file(entity: Entity):
async def write_file(entity: EntityIn):
await write_entity_file(self.entities_path, entity)
file_writes = [write_file(entity) for entity in entities]
file_writes = [write_file(entity) for entity in entities_in]
await asyncio.gather(*file_writes)
# Update database index sequentially
for entity in entities:
await self.entity_service.create_entity(entity)
async def create_entity_in_db(entity_in: EntityIn):
print(f"\nCreating entity in DB: {entity_in.name}")
db_entity = await self.entity_service.create_entity(entity_in)
print(f"Adding {len(entity_in.observations)} observations to DB for {entity_in.name}")
await self.observation_service.add_observations(entity_in, entity_in.observations)
[await self.relation_service.create_relation(relation_in) for relation_in in entity_in.relations]
# query the entity again to return relations
final_entity = await self.entity_service.get_entity(entity_in.id)
print(f"Final entity {final_entity.name} has {len(final_entity.observations)} observations in DB")
return final_entity
# Update database index sequentially
entities = [await create_entity_in_db(entities_in) for entities_in in entities_in]
return entities
async def create_relations(self, relations_data: List[Dict[str, Any]]) -> List[Relation]:
async def create_relations(self, relations_data: List[Dict[str, Any]]) -> List[RelationOut]:
"""Create multiple relations between entities."""
relations = [Relation.model_validate(data) for data in relations_data]
relations = [RelationIn.model_validate(data) for data in relations_data]
for relation in relations:
# First read complete entities from filesystem
@@ -68,46 +80,39 @@ class MemoryService:
return relations
async def add_observations(self, observations_in: Dict[str, Any]) -> ObservationsOut:
async def add_observations(self, observations_in: Dict[str, Any]) -> List[Observation]:
"""Add observations to an existing entity.
Args:
observations_in: input containing entity_name and observations
observations_in: input containing entity_id and observations
Returns:
ObservationsOut containing the created observations with IDs
List[Observation] with the newly created observations
"""
# Create new observations
new_observations = ObservationsIn.model_validate(observations_in)
print(f"\nAdding new observations to entity {new_observations.entity_id}")
print(f"New observations to add: {len(new_observations.observations)}")
# Read entity from filesystem
entity = await read_entity_file(self.entities_path, new_observations.entity_id)
# Convert ObservationIn to Observation before adding to entity
entity_observations = [
Observation(content=obs.content, context=obs.context)
for obs in new_observations.observations
]
entity.observations.extend(entity_observations)
print(f"Entity {entity.id} from file has {len(entity.observations)} observations")
# Create new observations for the entity
for obs in new_observations.observations:
entity.observations.append(obs)
print(f"After appending, entity has {len(entity.observations)} observations")
# Write updated entity file
await write_entity_file(self.entities_path, entity)
# Update database index
added_observations = await self.observation_service.add_observations(entity, new_observations.observations)
print(f"Added {len(added_observations)} observations to DB")
# Create and return output model
return ObservationsOut(
entity_id=entity.id,
observations=[
ObservationOut(
id=obs.id,
content=obs.content,
context=obs.context
)
for obs in added_observations
]
)
db_entity = await self.entity_service.get_entity(entity.id)
print(f"Entity {entity.id} in DB now has {len(db_entity.observations)} observations")
return added_observations
async def delete_entities(self, entity_names: List[str]) -> None:
"""Delete multiple entities and their associated data."""
@@ -1,15 +1,13 @@
"""Service for managing observations in both filesystem and database."""
from datetime import datetime, UTC
from pathlib import Path
from typing import Optional, List
from uuid import uuid4
from typing import List
from sqlalchemy import select, delete
from basic_memory.models import Observation as DbObservation
from basic_memory.models import Observation
from basic_memory.repository import ObservationRepository
from basic_memory.schemas import Entity, Observation, ObservationIn
from basic_memory.models import Observation as ObservationModel
from . import ServiceError, DatabaseSyncError
from basic_memory.schemas import EntityIn, ObservationIn
from . import DatabaseSyncError
class ObservationService:
@@ -22,38 +20,31 @@ class ObservationService:
self.project_path = project_path
self.observation_repo = observation_repo
async def add_observations(self, entity: Entity, observations: List[ObservationIn]) -> List[Observation]:
async def add_observations(self, entity: EntityIn, observations: List[ObservationIn]) -> List[Observation]:
"""
Add multiple observations to an entity.
Returns the created observations with IDs set.
"""
created_observations = []
print(f"\nObservationService.add_observations called for entity {entity.id}")
print(f"Adding {len(observations)} observations")
async def add_observation(observation: ObservationIn) -> Observation:
try:
db_observation = await self.observation_repo.create({
return await self.observation_repo.create({
'entity_id': entity.id,
'content': observation.content,
'context': observation.context,
'created_at': datetime.now(UTC)
})
# Convert db model to schema
return Observation(
id=db_observation.id,
content=db_observation.content,
context=db_observation.context
)
except Exception as e:
raise DatabaseSyncError(f"Failed to add observation to database: {str(e)}") from e
# Add each observation and collect the results
created_observations = [await add_observation(obs) for obs in observations]
# Update entity in memory with the created observations that have IDs
entity.observations.extend(created_observations)
print(f"Created {len(created_observations)} observations in DB")
return created_observations
async def search_observations(self, query: str) -> List[ObservationModel]:
async def search_observations(self, query: str) -> List[Observation]:
"""
Search for observations across all entities.
@@ -64,8 +55,8 @@ class ObservationService:
List of matching observations with their entity contexts
"""
result = await self.observation_repo.execute_query(
select(DbObservation).filter(
DbObservation.content.contains(query)
select(Observation).filter(
Observation.content.contains(query)
)
)
return [
@@ -73,29 +64,10 @@ class ObservationService:
for obs in result.scalars().all()
]
async def get_observations_by_context(self, context: str) -> List[ObservationModel]:
async def get_observations_by_context(self, context: str) -> List[Observation]:
"""Get all observations with a specific context."""
db_observations = await self.observation_repo.find_by_context(context)
return [
Observation(content=obs.content)
for obs in db_observations
]
async def rebuild_observation_index(self, entity: Entity) -> None:
"""
Rebuild the observation database index for a specific entity.
Used for recovery or ensuring sync.
"""
# Clear existing observations for this entity
await self.observation_repo.execute_query(
delete(DbObservation).where(DbObservation.entity_id == entity.id)
)
# Rebuild from entity's observations
for obs in entity.observations:
await self.observation_repo.create({
'entity_id': entity.id,
'content': obs.content,
'context': obs.context,
'created_at': datetime.now(UTC)
})
]
@@ -4,9 +4,9 @@ from pathlib import Path
from typing import Dict, Any
from sqlalchemy import delete
from basic_memory.models import Relation as DbRelation
from basic_memory.models import Relation as DbRelation, Relation
from basic_memory.repository import RelationRepository
from basic_memory.schemas import Entity, Relation
from basic_memory.schemas import EntityIn, RelationIn
from . import ServiceError, DatabaseSyncError, RelationError
@@ -20,7 +20,7 @@ class RelationService:
self.project_path = project_path
self.relation_repo = relation_repo
async def create_relation(self, relation: Relation) -> Relation:
async def create_relation(self, relation: RelationIn) -> Relation:
"""Create a new relation in the database."""
try:
db_data = relation.model_dump()
@@ -30,7 +30,7 @@ class RelationService:
except Exception as e:
raise DatabaseSyncError(f"Failed to sync relation to database: {str(e)}") from e
async def delete_relation(self, from_entity: Entity, to_entity: Entity, relation_type: str) -> bool:
async def delete_relation(self, from_entity: EntityIn, to_entity: EntityIn, relation_type: str) -> bool:
"""Delete a specific relation between entities."""
# Find and remove the relation from the entity's relations
if hasattr(from_entity, 'relations'):