| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157 |
- 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,
- )
|