From ff98c8068fc6381b980913451035dcc25706f782 Mon Sep 17 00:00:00 2001 From: phernandez Date: Thu, 12 Dec 2024 19:45:25 -0600 Subject: [PATCH] refactor observation create --- src/basic_memory/repository/__init__.py | 50 ++++++++++++++++--- src/basic_memory/services/memory_service.py | 12 +++-- .../services/observation_service.py | 46 ++++++----------- 3 files changed, 67 insertions(+), 41 deletions(-) diff --git a/src/basic_memory/repository/__init__.py b/src/basic_memory/repository/__init__.py index b4961fd3..ee6481ea 100644 --- a/src/basic_memory/repository/__init__.py +++ b/src/basic_memory/repository/__init__.py @@ -1,5 +1,5 @@ """Base repository implementation.""" -from typing import Type, Optional, Any, Sequence, TypeVar +from typing import Type, Optional, Any, Sequence, TypeVar, List from sqlalchemy import select, func, Select, Executable, inspect, Result, Column, insert from sqlalchemy.exc import NoResultFound from sqlalchemy.ext.asyncio import AsyncSession @@ -69,17 +69,55 @@ class Repository[T: Base]: model = model or self.Model logger.debug(f"Creating {model.__name__} with data: {entity_data}") try: + # Only include valid columns that are provided in entity_data + model_data = { + k: v for k, v in entity_data.items() + if k in self.valid_columns and v is not None + } + + # Generate ID if this is an Entity model and no ID provided + if model is Entity and 'id' not in model_data: + model_data['id'] = Entity.generate_id( + model_data['entity_type'], + model_data['name'] + ) + + logger.debug(f"Filtered data for valid columns: {model_data}") # Create insert statement with only provided data - instance = model(**entity_data) - self.session.add(instance) - await self.session.flush() - logger.debug(f"Created {model.__name__}: {getattr(instance, 'id', None)}") - return instance + stmt = insert(model).values(**model_data).returning(model) + result = await self.session.execute(stmt) + entity = result.scalar_one() + + logger.debug(f"Created {model.__name__}: {getattr(entity, 'id', None)}") + return entity except Exception as e: logger.exception(f"Failed to create {model.__name__}") raise + + async def instance_create(self, instance: T) -> T: + """Create a new record from a model instance.""" + logger.debug(f"Creating {self.Model.__name__} from instance: {instance}") + try: + self.session.add(instance) + await self.session.flush() + return instance + except Exception as e: + logger.exception(f"Failed to create {self.Model.__name__}") + raise + + async def bulk_create(self, instances: List[T]) -> List[T]: + """Create multiple records in a single transaction.""" + logger.debug(f"Bulk creating {len(instances)} {self.Model.__name__} instances") + try: + for instance in instances: + self.session.add(instance) + await self.session.flush() + return instances + except Exception as e: + logger.exception(f"Failed to bulk create {self.Model.__name__}") + raise async def update(self, entity_id: str, entity_data: dict) -> Optional[T]: """Update an entity with the given data.""" diff --git a/src/basic_memory/services/memory_service.py b/src/basic_memory/services/memory_service.py index af958037..a876de7e 100644 --- a/src/basic_memory/services/memory_service.py +++ b/src/basic_memory/services/memory_service.py @@ -5,7 +5,7 @@ from pathlib import Path from basic_memory.models import Entity, Observation, Relation from basic_memory.schemas import ( - ObservationsIn, EntityIn, RelationIn + ObservationsIn, EntityIn, RelationIn, ObservationIn ) from basic_memory.fileio import write_entity_file, read_entity_file, EntityNotFoundError from basic_memory.services import EntityService, RelationService, ObservationService @@ -64,9 +64,13 @@ class MemoryService: created_entity = await self.entity_service.create_entity(entity_in) logger.debug(f"Created base entity: {created_entity.id}") - # Add observations - await self.observation_service.add_observations(created_entity.id, entity_in.observations) - logger.debug(f"Added {len(entity_in.observations)} observations to {created_entity.id}") + # Convert ObservationIn to Observation instances + if entity_in.observations: + created_observations = await self.observation_service.add_observations( + created_entity.id, + [ObservationIn(**obs.model_dump()) for obs in entity_in.observations] + ) + logger.debug(f"Added {len(created_observations)} observations to {created_entity.id}") # Add relations for relation in entity_in.relations: diff --git a/src/basic_memory/services/observation_service.py b/src/basic_memory/services/observation_service.py index 0a99e39c..e4194b17 100644 --- a/src/basic_memory/services/observation_service.py +++ b/src/basic_memory/services/observation_service.py @@ -1,11 +1,11 @@ """Service for managing observations in both filesystem and database.""" from pathlib import Path -from typing import List +from typing import List, Sequence from sqlalchemy import select from basic_memory.models import Observation from basic_memory.repository.observation_repository import ObservationRepository -from basic_memory.schemas import EntityIn, ObservationIn +from basic_memory.schemas import ObservationIn from . import DatabaseSyncError @@ -24,29 +24,17 @@ class ObservationService: Add multiple observations to an entity. Returns the created observations with IDs set. """ - async def add_observation(observation: ObservationIn) -> Observation: - try: - obs = await self.observation_repo.create({ - **observation.model_dump(), - 'entity_id': entity_id - }) - # Ensure observation is flushed and refreshed - await self.observation_repo.session.flush() - await self.observation_repo.session.refresh(obs) - return obs - 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] - - # Make sure observations are in sync before returning - # This helps ensure related entities see the new observations - await self.observation_repo.session.flush() - for obs in created_observations: - await self.observation_repo.session.refresh(obs) - - return created_observations + try: + return await self.observation_repo.bulk_create([ + Observation( + entity_id=entity_id, + content=observation.content, + context=observation.context + ) + for observation in observations + ]) + except Exception as e: + raise DatabaseSyncError(f"Failed to add observations to database: {str(e)}") from e async def search_observations(self, query: str) -> List[Observation]: """ @@ -68,10 +56,6 @@ class ObservationService: for obs in result.scalars().all() ] - async def get_observations_by_context(self, context: str) -> List[Observation]: + async def get_observations_by_context(self, context: str) -> Sequence[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 - ] \ No newline at end of file + return await self.observation_repo.find_by_context(context)