agent_repositories.py 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157
  1. from sqlalchemy import select
  2. from app.domains.agent.models import AgentMessage, AgentRun, AgentThread
  3. from app.domains.agent.repository import AgentRepository
  4. from app.infrastructure.mysql.agent_models import (
  5. AgentMessageRecord,
  6. AgentRunRecord,
  7. AgentThreadRecord,
  8. )
  9. from app.infrastructure.mysql.repositories import SessionFactory
  10. class SqlAlchemyAgentRepository(AgentRepository):
  11. def __init__(self, session_factory: SessionFactory) -> None:
  12. self._session_factory = session_factory
  13. def save_thread(self, thread: AgentThread) -> None:
  14. with self._session_factory() as session:
  15. session.merge(
  16. AgentThreadRecord(
  17. id=thread.id,
  18. owner_type=thread.owner_type,
  19. owner_id=thread.owner_id,
  20. persona=thread.persona,
  21. title=thread.title,
  22. status=thread.status,
  23. created_at=thread.created_at,
  24. updated_at=thread.updated_at,
  25. )
  26. )
  27. session.commit()
  28. def get_thread(self, thread_id: str) -> AgentThread | None:
  29. with self._session_factory() as session:
  30. record = session.get(AgentThreadRecord, thread_id)
  31. return self._to_domain(record) if record else None
  32. def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]:
  33. with self._session_factory() as session:
  34. records = session.scalars(
  35. select(AgentThreadRecord).where(
  36. AgentThreadRecord.owner_type == owner_type,
  37. AgentThreadRecord.owner_id == owner_id,
  38. )
  39. ).all()
  40. return [self._to_domain(record) for record in records]
  41. def save_message(self, message: AgentMessage) -> None:
  42. with self._session_factory() as session:
  43. session.merge(
  44. AgentMessageRecord(
  45. id=message.id,
  46. thread_id=message.thread_id,
  47. role=message.role,
  48. content=message.content,
  49. metadata_json=message.metadata,
  50. created_at=message.created_at,
  51. )
  52. )
  53. session.commit()
  54. def list_messages(self, thread_id: str) -> list[AgentMessage]:
  55. with self._session_factory() as session:
  56. records = session.scalars(
  57. select(AgentMessageRecord)
  58. .where(AgentMessageRecord.thread_id == thread_id)
  59. .order_by(AgentMessageRecord.created_at)
  60. ).all()
  61. return [
  62. AgentMessage(
  63. id=record.id,
  64. thread_id=record.thread_id,
  65. role=record.role,
  66. content=record.content,
  67. metadata=dict(record.metadata_json),
  68. created_at=record.created_at,
  69. )
  70. for record in records
  71. ]
  72. def save_run(self, run: AgentRun) -> None:
  73. with self._session_factory() as session:
  74. session.merge(
  75. AgentRunRecord(
  76. id=run.id,
  77. request_id=run.request_id,
  78. thread_id=run.thread_id,
  79. trace_id=run.trace_id,
  80. persona=run.persona,
  81. status=run.status,
  82. input_json=run.input,
  83. output_json=run.output,
  84. error_code=run.error_code,
  85. created_at=run.created_at,
  86. completed_at=run.completed_at,
  87. )
  88. )
  89. session.commit()
  90. def get_run(self, run_id: str) -> AgentRun | None:
  91. with self._session_factory() as session:
  92. record = session.get(AgentRunRecord, run_id)
  93. if record is None:
  94. return None
  95. return AgentRun(
  96. id=record.id,
  97. request_id=record.request_id,
  98. thread_id=record.thread_id,
  99. trace_id=record.trace_id,
  100. persona=record.persona,
  101. status=record.status,
  102. input=dict(record.input_json),
  103. output=dict(record.output_json) if record.output_json else None,
  104. error_code=record.error_code,
  105. created_at=record.created_at,
  106. completed_at=record.completed_at,
  107. )
  108. def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]:
  109. if not thread_ids:
  110. return []
  111. with self._session_factory() as session:
  112. records = session.scalars(
  113. select(AgentRunRecord)
  114. .where(AgentRunRecord.thread_id.in_(thread_ids))
  115. .order_by(AgentRunRecord.created_at.desc())
  116. .limit(limit)
  117. ).all()
  118. return [
  119. AgentRun(
  120. id=record.id,
  121. request_id=record.request_id,
  122. thread_id=record.thread_id,
  123. trace_id=record.trace_id,
  124. persona=record.persona,
  125. status=record.status,
  126. input=dict(record.input_json),
  127. output=dict(record.output_json) if record.output_json else None,
  128. error_code=record.error_code,
  129. created_at=record.created_at,
  130. completed_at=record.completed_at,
  131. )
  132. for record in records
  133. ]
  134. @staticmethod
  135. def _to_domain(record: AgentThreadRecord) -> AgentThread:
  136. return AgentThread(
  137. id=record.id,
  138. owner_type=record.owner_type,
  139. owner_id=record.owner_id,
  140. persona=record.persona,
  141. title=record.title,
  142. status=record.status,
  143. created_at=record.created_at,
  144. updated_at=record.updated_at,
  145. )