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]