from sqlalchemy import select from app.domains.agent.models import AgentMessage, AgentRun, AgentThread from app.domains.agent.repository import AgentRepository from app.infrastructure.mysql.agent_models import ( AgentMessageRecord, AgentRunRecord, AgentThreadRecord, ) from app.infrastructure.mysql.repositories import SessionFactory class SqlAlchemyAgentRepository(AgentRepository): def __init__(self, session_factory: SessionFactory) -> None: self._session_factory = session_factory def save_thread(self, thread: AgentThread) -> None: with self._session_factory() as session: session.merge( AgentThreadRecord( id=thread.id, owner_type=thread.owner_type, owner_id=thread.owner_id, persona=thread.persona, title=thread.title, status=thread.status, created_at=thread.created_at, updated_at=thread.updated_at, ) ) session.commit() def get_thread(self, thread_id: str) -> AgentThread | None: with self._session_factory() as session: record = session.get(AgentThreadRecord, thread_id) return self._to_domain(record) if record else None def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]: with self._session_factory() as session: records = session.scalars( select(AgentThreadRecord).where( AgentThreadRecord.owner_type == owner_type, AgentThreadRecord.owner_id == owner_id, ) ).all() return [self._to_domain(record) for record in records] def save_message(self, message: AgentMessage) -> None: with self._session_factory() as session: session.merge( AgentMessageRecord( id=message.id, thread_id=message.thread_id, role=message.role, content=message.content, metadata_json=message.metadata, created_at=message.created_at, ) ) session.commit() def list_messages(self, thread_id: str) -> list[AgentMessage]: with self._session_factory() as session: records = session.scalars( select(AgentMessageRecord) .where(AgentMessageRecord.thread_id == thread_id) .order_by(AgentMessageRecord.created_at) ).all() return [ AgentMessage( id=record.id, thread_id=record.thread_id, role=record.role, content=record.content, metadata=dict(record.metadata_json), created_at=record.created_at, ) for record in records ] def save_run(self, run: AgentRun) -> None: with self._session_factory() as session: session.merge( AgentRunRecord( id=run.id, request_id=run.request_id, thread_id=run.thread_id, trace_id=run.trace_id, persona=run.persona, status=run.status, input_json=run.input, output_json=run.output, error_code=run.error_code, created_at=run.created_at, completed_at=run.completed_at, ) ) session.commit() def get_run(self, run_id: str) -> AgentRun | None: with self._session_factory() as session: record = session.get(AgentRunRecord, run_id) if record is None: return None return AgentRun( id=record.id, request_id=record.request_id, thread_id=record.thread_id, trace_id=record.trace_id, persona=record.persona, status=record.status, input=dict(record.input_json), output=dict(record.output_json) if record.output_json else None, error_code=record.error_code, created_at=record.created_at, completed_at=record.completed_at, ) def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]: if not thread_ids: return [] with self._session_factory() as session: records = session.scalars( select(AgentRunRecord) .where(AgentRunRecord.thread_id.in_(thread_ids)) .order_by(AgentRunRecord.created_at.desc()) .limit(limit) ).all() return [ AgentRun( id=record.id, request_id=record.request_id, thread_id=record.thread_id, trace_id=record.trace_id, persona=record.persona, status=record.status, input=dict(record.input_json), output=dict(record.output_json) if record.output_json else None, error_code=record.error_code, created_at=record.created_at, completed_at=record.completed_at, ) for record in records ] @staticmethod def _to_domain(record: AgentThreadRecord) -> AgentThread: return AgentThread( id=record.id, owner_type=record.owner_type, owner_id=record.owner_id, persona=record.persona, title=record.title, status=record.status, created_at=record.created_at, updated_at=record.updated_at, )