[hAIk]: refactor AgentCard and BrowserCard into subdirectory modules, plumb dashboard_id from launch through Agent/AgentSnapshot, add on_done auto-persist callback, fix get_all_sessions to use GET query params, atomic PydanticStore writes, graceful backend shutdown in local.sh, and add mode config JSONs

This commit is contained in:
haikdc
2026-04-18 12:05:50 -07:00
parent 21a723a45b
commit 6a9dd59eec
22 changed files with 139 additions and 91 deletions
+22 -3
View File
@@ -48,6 +48,14 @@ AGENT_STORE: PydanticStore[Agent] = PydanticStore[Agent](
# NOTE: Essentially the SESSIONS is a cache for active agents.
SESSIONS: dict[str, Agent] = {}
def _persist_agent(agent: Agent) -> None:
"""Called when an agent reaches a terminal state (completed/error)."""
debug("auto-saving session %s (status=%s)", agent.session_id, agent.status)
try:
AGENT_STORE.save(agent)
except Exception as e:
debug("auto-save failed for session %s: %s", agent.session_id, e)
def get_agent(session_id: str) -> Agent:
agent: Optional[Agent] = SESSIONS.get(session_id)
if not agent:
@@ -64,6 +72,7 @@ async def agents_lifespan():
try:
stored.status = "stopped"
stored.on_event = COMMS_MANAGER.make_session_emitter(stored.session_id)
stored.on_done = _persist_agent
stored.toolkit = await build_agent_toolkit(
agent=stored,
sessions=SESSIONS,
@@ -73,10 +82,15 @@ async def agents_lifespan():
except Exception as e:
debug(f"[agents lifespan] Skipping corrupt session {stored.session_id}: {e}")
yield
debug("agents_lifespan: shutting down — %s sessions to save", len(SESSIONS))
for agent in list[Agent](SESSIONS.values()):
debug("agents_lifespan: stopping agent %s", agent.session_id)
await agent.stop_agent()
debug("agents_lifespan: saving agent %s", agent.session_id)
AGENT_STORE.save(agent)
debug("agents_lifespan: saved agent %s", agent.session_id)
SESSIONS.clear()
debug("agents_lifespan: shutdown complete")
agents = SubApp("agents", agents_lifespan)
@@ -115,10 +129,10 @@ async def websocket_dashboard(websocket: WebSocket):
# ---------------------------------------------------------------------------
@agents.router.get("/get_all_sessions")
async def get_all_sessions(dashboard_id: str = Body(default="")) -> dict:
result: List[Agent] = list[Agent](SESSIONS.values())
async def get_all_sessions(dashboard_id: str = "") -> dict:
result: List[Agent] = list(SESSIONS.values())
if dashboard_id:
result: List[Agent] = [a for a in result if getattr(a, "dashboard_id", None) == dashboard_id]
result = [a for a in result if a.dashboard_id == dashboard_id]
return {"sessions": [a.model_dump(mode="json") for a in result]}
@@ -134,14 +148,17 @@ async def launch_agent(
mode: str = Body(),
system_prompt: str = Body(),
max_turns: int = Body(),
dashboard_id: Optional[str] = Body(default=None),
) -> dict:
agent: Agent = Agent(
model=model,
mode=mode,
status="stopped",
dashboard_id=dashboard_id,
config=ClaudeAgentOptions(max_turns=max_turns),
)
agent.on_event = COMMS_MANAGER.make_session_emitter(agent.session_id)
agent.on_done = _persist_agent
SESSIONS[agent.session_id] = agent
toolkit: Toolkit = await build_agent_toolkit(
@@ -372,6 +389,7 @@ async def resume_session(session_id: str = Body()) -> dict:
raise HTTPException(status_code=404, detail="Session not found in history")
agent.status = "stopped"
agent.on_event = COMMS_MANAGER.make_session_emitter(agent.session_id)
agent.on_done = _persist_agent
agent.toolkit = await build_agent_toolkit(
agent=agent,
sessions=SESSIONS,
@@ -402,6 +420,7 @@ async def duplicate_session(session_id: str = Body()) -> dict:
clone.pending_approvals = []
clone.sub_agents = []
clone.on_event = COMMS_MANAGER.make_session_emitter(clone.session_id)
clone.on_done = _persist_agent
clone.toolkit = await build_agent_toolkit(
agent=clone,
sessions=SESSIONS,
+9 -3
View File
@@ -32,7 +32,8 @@ class Agent(BaseModel):
messages: MessageLog = Field(default_factory=MessageLog)
session_id: str = Field(default_factory=lambda: uuid4().hex)
config: ClaudeAgentOptions
dashboard_id: Optional[str] = None
config: ClaudeAgentOptions = Field(default_factory=ClaudeAgentOptions, exclude=True)
branch_id: str = "main"
sub_agents: List["Agent"] = Field(default_factory=list)
@@ -41,9 +42,10 @@ class Agent(BaseModel):
toolkit: Optional[Toolkit] = Field(default=None, exclude=True)
on_event: Optional[EventCallback] = Field(default=None, exclude=True)
on_done: Optional[Any] = Field(default=None, exclude=True)
task: Optional[InstanceOf[asyncio.Task]] = None
lock: InstanceOf[asyncio.Lock] = Field(default_factory=asyncio.Lock)
task: Optional[InstanceOf[asyncio.Task]] = Field(default=None, exclude=True)
lock: InstanceOf[asyncio.Lock] = Field(default_factory=asyncio.Lock, exclude=True)
@typechecked
def snapshot(self) -> AgentSnapshot:
@@ -52,6 +54,7 @@ class Agent(BaseModel):
model=self.model,
mode=self.mode,
status=self.status,
dashboard_id=self.dashboard_id,
branch_id=self.branch_id,
parent_id=self.parent_id,
messages=self.messages,
@@ -70,6 +73,9 @@ class Agent(BaseModel):
self.status = event.status # type: ignore[assignment]
if self.on_event:
await self.on_event(event)
if isinstance(event, AgentStatusEvent) and event.status in ("completed", "error"):
if self.on_done:
self.on_done(self)
@typechecked
async def request_approval(
+22 -6
View File
@@ -6,10 +6,12 @@ inside a data directory. This module eliminates that copy-paste.
import json
import os
import tempfile
from typing import Generic, List, Optional, TypeVar
from fastapi import HTTPException
from pydantic import BaseModel
from swarm_debug import debug
from typeguard import typechecked
T = TypeVar("T", bound=BaseModel)
@@ -39,19 +41,33 @@ class PydanticStore(BaseModel, Generic[T]):
def load_all(self) -> list[T]:
result: List[T] = []
if not os.path.exists(self.data_dir):
debug("load_all: data_dir does not exist: %s", self.data_dir)
return result
for fname in os.listdir(self.data_dir):
if fname.endswith(".json"):
with open(os.path.join(self.data_dir, fname)) as f:
result.append(self.model_cls(**json.load(f)))
fnames = [f for f in os.listdir(self.data_dir) if f.endswith(".json")]
debug("load_all: scanning %s — found %s json files", self.data_dir, len(fnames), table=False)
for fname in fnames:
path = os.path.join(self.data_dir, fname)
size = os.path.getsize(path)
debug("load_all: loading %s (%s bytes)", fname, size, table=False)
with open(path) as f:
raw = f.read()
debug("load_all: raw content length=%s, first 100 chars: %s", len(raw), raw[:100], table=False)
result.append(self.model_cls(**json.loads(raw)))
return result
@typechecked
def save(self, item: T) -> None:
os.makedirs(self.data_dir, exist_ok=True)
item_id = getattr(item, self.id_field)
with open(self.p_path(item_id), "w") as f:
json.dump(self.p_dump(item), f, indent=2)
path = self.p_path(item_id)
debug("save: writing %s to %s", item_id, path, table=False)
with tempfile.NamedTemporaryFile(
"w", dir=self.data_dir, suffix=".tmp", delete=False
) as tmp:
json.dump(self.p_dump(item), tmp, indent=2)
tmp_path = tmp.name
os.replace(tmp_path, path)
debug("save: complete — %s is now %s bytes", path, os.path.getsize(path), table=False)
@typechecked
def load(self, item_id: str) -> T:
@@ -9,6 +9,7 @@ class AgentSnapshot(BaseModel):
model: str
mode: str
status: str
dashboard_id: Optional[str] = None
branch_id: str = "main"
parent_id: Optional[str] = None
messages: MessageLog = Field(default_factory=MessageLog)