repository.py 2.1 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758
  1. from typing import Protocol
  2. from app.domains.agent.models import AgentMessage, AgentRun, AgentThread
  3. class AgentRepository(Protocol):
  4. def save_thread(self, thread: AgentThread) -> None: ...
  5. def get_thread(self, thread_id: str) -> AgentThread | None: ...
  6. def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]: ...
  7. def save_message(self, message: AgentMessage) -> None: ...
  8. def list_messages(self, thread_id: str) -> list[AgentMessage]: ...
  9. def save_run(self, run: AgentRun) -> None: ...
  10. def get_run(self, run_id: str) -> AgentRun | None: ...
  11. def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]: ...
  12. class InMemoryAgentRepository:
  13. def __init__(self) -> None:
  14. self._threads: dict[str, AgentThread] = {}
  15. self._messages: dict[str, AgentMessage] = {}
  16. self._runs: dict[str, AgentRun] = {}
  17. def save_thread(self, thread: AgentThread) -> None:
  18. self._threads[thread.id] = thread
  19. def get_thread(self, thread_id: str) -> AgentThread | None:
  20. return self._threads.get(thread_id)
  21. def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]:
  22. return [
  23. thread
  24. for thread in self._threads.values()
  25. if thread.owner_type == owner_type and thread.owner_id == owner_id
  26. ]
  27. def save_message(self, message: AgentMessage) -> None:
  28. self._messages[message.id] = message
  29. def list_messages(self, thread_id: str) -> list[AgentMessage]:
  30. return [message for message in self._messages.values() if message.thread_id == thread_id]
  31. def save_run(self, run: AgentRun) -> None:
  32. self._runs[run.id] = run
  33. def get_run(self, run_id: str) -> AgentRun | None:
  34. return self._runs.get(run_id)
  35. def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]:
  36. allowed = set(thread_ids)
  37. runs = [run for run in self._runs.values() if run.thread_id in allowed]
  38. return sorted(runs, key=lambda run: run.created_at, reverse=True)[:limit]