| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263 |
- from typing import Protocol
- from zbt.domains.agent.models import AgentMessage, AgentRun, AgentThread
- class AgentRepository(Protocol):
- def save_thread(self, thread: AgentThread) -> None: ...
- def get_thread(self, thread_id: str) -> AgentThread | None: ...
- def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]: ...
- def save_message(self, message: AgentMessage) -> None: ...
- def list_messages(self, thread_id: str) -> list[AgentMessage]: ...
- def save_run(self, run: AgentRun) -> None: ...
- def get_run(self, run_id: str) -> AgentRun | None: ...
- def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]: ...
- class InMemoryAgentRepository:
- def __init__(self) -> None:
- self._threads: dict[str, AgentThread] = {}
- self._messages: dict[str, AgentMessage] = {}
- self._runs: dict[str, AgentRun] = {}
- def save_thread(self, thread: AgentThread) -> None:
- self._threads[thread.id] = thread
- def get_thread(self, thread_id: str) -> AgentThread | None:
- return self._threads.get(thread_id)
- def list_threads(self, owner_type: str, owner_id: str) -> list[AgentThread]:
- return [
- thread
- for thread in self._threads.values()
- if thread.owner_type == owner_type and thread.owner_id == owner_id
- ]
- def save_message(self, message: AgentMessage) -> None:
- self._messages[message.id] = message
- def list_messages(self, thread_id: str) -> list[AgentMessage]:
- messages = [
- message
- for message in self._messages.values()
- if message.thread_id == thread_id
- ]
- return sorted(messages, key=lambda message: message.created_at)
- def save_run(self, run: AgentRun) -> None:
- self._runs[run.id] = run
- def get_run(self, run_id: str) -> AgentRun | None:
- return self._runs.get(run_id)
- def list_runs(self, thread_ids: tuple[str, ...], limit: int = 20) -> list[AgentRun]:
- allowed = set(thread_ids)
- runs = [run for run in self._runs.values() if run.thread_id in allowed]
- return sorted(runs, key=lambda run: run.created_at, reverse=True)[:limit]
|