default created_at for observation, relation

This commit is contained in:
phernandez
2024-12-12 20:29:11 -06:00
parent ff98c8068f
commit d56197c288
4 changed files with 84 additions and 47 deletions
+2 -13
View File
@@ -4,7 +4,6 @@ from typing import List, Optional
from sqlalchemy import String, DateTime, ForeignKey, Text, TypeDecorator, Integer, text, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship, DeclarativeBase
from sqlalchemy.ext.asyncio import AsyncAttrs
from sqlalchemy import orm
class UTCDateTime(TypeDecorator):
@@ -29,20 +28,10 @@ def utc_now() -> datetime:
"""Helper to get current UTC time"""
return datetime.now(UTC)
def lenient_constructor(self, **kwargs):
cls_ = type(self)
for k in kwargs:
if not hasattr(cls_, k):
print(f'Skipping invalid attr {k!r}')
continue
setattr(self, k, kwargs[k])
registry = orm.registry(constructor=lenient_constructor)
class Base(AsyncAttrs, DeclarativeBase):
"""Base class for all models"""
registry = registry
pass
class Entity(Base):
@@ -127,7 +116,7 @@ class Observation(Base):
)
content: Mapped[str] = mapped_column(String)
created_at: Mapped[datetime] = mapped_column(
UTCDateTime,
DateTime,
server_default=text('CURRENT_TIMESTAMP')
)
context: Mapped[str | None] = mapped_column(String, nullable=True)
+9 -10
View File
@@ -22,8 +22,11 @@ class MemoryService:
relation_service: RelationService,
observation_service: ObservationService
):
self.project_path = project_path
self.entities_path = project_path / "entities" if project_path else None
if project_path:
assert project_path.is_dir(), "Path does not exist or is not a directory: {project_path}"
self.project_path = project_path
self.entities_path = project_path / "entities"
self.entity_service = entity_service
self.relation_service = relation_service
self.observation_service = observation_service
@@ -64,13 +67,9 @@ class MemoryService:
created_entity = await self.entity_service.create_entity(entity_in)
logger.debug(f"Created base entity: {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 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}")
# Add relations
for relation in entity_in.relations:
@@ -104,7 +103,7 @@ class MemoryService:
if path.exists():
path.unlink()
except Exception as cleanup_error:
logger.error(f"Failed to clean up file for {entity_id}: {cleanup_error}")
logger.error(f"Failed to clean up file for {entity.id}: {cleanup_error}")
raise
async def create_relations(self, relations_data: List[RelationIn]) -> List[Relation]: