mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-30 21:44:50 +02:00
[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:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user