mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-11 12:17:45 +02:00
[Haik]: ckpt, largely done w the initial Agent class refactor. Now gonna add sub agent references for recursive cancels, then im gonna clean up the Agent folder for better downward abstraction
This commit is contained in:
@@ -22,13 +22,14 @@ from pydantic import BaseModel, InstanceOf
|
||||
from typing import List, Literal, Optional
|
||||
from typeguard import typechecked
|
||||
|
||||
from backend.apps.agents.models import AgentSession, Message
|
||||
from backend.apps.agents.manager.AgentConfig import AgentConfig
|
||||
from backend.apps.agents.models import AgentSession, AgentConfig, ApprovalRequest
|
||||
from backend.apps.agents.manager.ws_manager import ws_manager
|
||||
from backend.apps.agents.execution.prompt_builder import resolve_mode
|
||||
from backend.apps.agents.execution.mcp_builder import get_all_tool_names
|
||||
from backend.apps.agents.execution.agent_loop import run_agent_loop
|
||||
from backend.apps.settings.settings import load_settings
|
||||
from backend.apps.agents.manager.HaikFix.agent_loop import run_agent_loop
|
||||
from backend.apps.agents.manager.HaikFix.helpers.Message import Message
|
||||
from backend.apps.agents.manager.HaikFix.PromptChunks import ImageChunk
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -51,6 +52,9 @@ class Agent(BaseModel):
|
||||
status: Literal["running", "waiting_approval", "completed", "error", "stopped"]
|
||||
lock: InstanceOf[asyncio.Lock]
|
||||
config: ClaudeAgentOptions
|
||||
session: Optional[AgentSession] = None
|
||||
branch_id: str = "main"
|
||||
parent_id: Optional[str] = None
|
||||
task: Optional[asyncio.Task] = None
|
||||
|
||||
@typechecked
|
||||
@@ -87,7 +91,7 @@ class Agent(BaseModel):
|
||||
if config.mode in ("view-builder", "skill-builder") and not config.target_directory:
|
||||
effective_cwd = os.path.join(effective_cwd, session_id)
|
||||
os.makedirs(effective_cwd, exist_ok=True)
|
||||
session = AgentSession(
|
||||
self.session = AgentSession(
|
||||
id=session_id, name=config.name,
|
||||
provider=getattr(config, "provider", "anthropic"),
|
||||
model=config.model, mode=config.mode,
|
||||
@@ -95,103 +99,55 @@ class Agent(BaseModel):
|
||||
max_turns=config.max_turns, cwd=effective_cwd,
|
||||
dashboard_id=config.dashboard_id,
|
||||
)
|
||||
await ws_manager.emit_status(session_id, "running", session)
|
||||
return session
|
||||
self.session_id = session_id
|
||||
await ws_manager.emit_status(session_id, "running", self.session)
|
||||
return self.session
|
||||
|
||||
async def send_message(
|
||||
self,
|
||||
prompt: str,
|
||||
images: Optional[list] = None,
|
||||
images: Optional[List[ImageChunk]] = None,
|
||||
):
|
||||
async with self.lock:
|
||||
if self.task is not None and not self.task.done():
|
||||
print("[Agent.send_message] Agent is already running")
|
||||
return
|
||||
|
||||
user_msg = Message(
|
||||
role="user",
|
||||
content=prompt,
|
||||
branch_id=self.branch_id,
|
||||
parent_id=self.parent_id,
|
||||
images=images,
|
||||
)
|
||||
|
||||
skill_meta = [{"id": s["id"], "name": s["name"]} for s in (attached_skills or [])] or None
|
||||
image_meta = [{"data": img["data"], "media_type": img.get("media_type", "image/png")} for img in (images or [])] or None
|
||||
user_msg = Message(
|
||||
role="user", content=prompt, branch_id=session.active_branch_id,
|
||||
context_paths=context_paths or None, attached_skills=skill_meta,
|
||||
forced_tools=forced_tools or None, images=image_meta, hidden=hidden,
|
||||
)
|
||||
session.messages.append(user_msg)
|
||||
await ws_manager.emit_message(session_id, user_msg)
|
||||
await ws_manager.emit_message(self.session_id, user_msg)
|
||||
self.status = "running"
|
||||
await ws_manager.emit_status(self.session_id, "running", self)
|
||||
|
||||
session.status = "running"
|
||||
await ws_manager.emit_status(session_id, "running", session)
|
||||
task = asyncio.create_task(run_agent_loop(
|
||||
self.sessions, session_id, prompt, images=images,
|
||||
context_paths=context_paths, forced_tools=forced_tools,
|
||||
attached_skills=attached_skills, selected_browser_ids=selected_browser_ids,
|
||||
))
|
||||
self.tasks[session_id] = task
|
||||
|
||||
async def send_message_old(
|
||||
self, session_id: str, prompt: str,
|
||||
mode: str | None = None, model: str | None = None,
|
||||
provider: str | None = None, images: list | None = None,
|
||||
context_paths: list | None = None, forced_tools: list[str] | None = None,
|
||||
attached_skills: list | None = None, hidden: bool = False,
|
||||
selected_browser_ids: list[str] | None = None,
|
||||
):
|
||||
session = self.sessions.get(session_id)
|
||||
if not session:
|
||||
raise ValueError(f"Session {session_id} not found")
|
||||
existing = self.tasks.get(session_id)
|
||||
if existing and not existing.done():
|
||||
return
|
||||
self.task = asyncio.create_task(run_agent_loop(
|
||||
prompt=prompt,
|
||||
images=images,
|
||||
options=self.config,
|
||||
branch_id=self.branch_id,
|
||||
parent_id=self.parent_id,
|
||||
))
|
||||
|
||||
session_changed = False
|
||||
if model and model != session.model:
|
||||
session.model = model
|
||||
session_changed = True
|
||||
if mode and mode != session.mode:
|
||||
session.mode = mode
|
||||
mode_tools, _, _ = resolve_mode(mode, get_all_tool_names)
|
||||
session.allowed_tools = mode_tools
|
||||
session_changed = True
|
||||
if session_changed:
|
||||
await ws_manager.emit_status(session_id, session.status, session)
|
||||
|
||||
skill_meta = [{"id": s["id"], "name": s["name"]} for s in (attached_skills or [])] or None
|
||||
image_meta = [{"data": img["data"], "media_type": img.get("media_type", "image/png")} for img in (images or [])] or None
|
||||
user_msg = Message(
|
||||
role="user", content=prompt, branch_id=session.active_branch_id,
|
||||
context_paths=context_paths or None, attached_skills=skill_meta,
|
||||
forced_tools=forced_tools or None, images=image_meta, hidden=hidden,
|
||||
)
|
||||
session.messages.append(user_msg)
|
||||
await ws_manager.emit_message(session_id, user_msg)
|
||||
|
||||
session.status = "running"
|
||||
await ws_manager.emit_status(session_id, "running", session)
|
||||
task = asyncio.create_task(run_agent_loop(
|
||||
self.sessions, session_id, prompt, images=images,
|
||||
context_paths=context_paths, forced_tools=forced_tools,
|
||||
attached_skills=attached_skills, selected_browser_ids=selected_browser_ids,
|
||||
))
|
||||
self.tasks[session_id] = task
|
||||
|
||||
async def stop_agent(self, session_id: str):
|
||||
task = self.tasks.get(session_id)
|
||||
if task and not task.done():
|
||||
task.cancel()
|
||||
async def stop_agent(self):
|
||||
if self.task and not self.task.done():
|
||||
self.task.cancel()
|
||||
try:
|
||||
await task
|
||||
await self.task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
session = self.sessions.get(session_id)
|
||||
if session:
|
||||
for req in list(session.pending_approvals):
|
||||
if self.session:
|
||||
for req in list[ApprovalRequest](self.session.pending_approvals):
|
||||
ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Agent stopped"})
|
||||
session.pending_approvals = []
|
||||
if hasattr(session, '_cancel_event'):
|
||||
session._cancel_event.set()
|
||||
session.status = "stopped"
|
||||
if not session.closed_at:
|
||||
session.closed_at = datetime.now()
|
||||
await ws_manager.emit_status(session_id, "stopped", session)
|
||||
children = [s for s in self.sessions.values() if s.parent_session_id == session_id and s.mode == "browser-agent"]
|
||||
for child in children:
|
||||
await self.stop_agent(child.id)
|
||||
self.session.pending_approvals = []
|
||||
if hasattr(self.session, '_cancel_event'):
|
||||
self.session._cancel_event.set()
|
||||
self.session.status = "stopped"
|
||||
if not self.session.closed_at:
|
||||
self.session.closed_at = datetime.now()
|
||||
await ws_manager.emit_status(self.session.id, "stopped", self.session)
|
||||
self.status = "stopped"
|
||||
|
||||
@@ -7,35 +7,18 @@ Heavy logic is delegated to sibling modules:
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from uuid import uuid4
|
||||
|
||||
from typeguard import typechecked
|
||||
|
||||
from backend.apps.agents.models import AgentSession, Message
|
||||
from backend.apps.agents.manager.ws_manager import ws_manager
|
||||
from backend.apps.agents.manager.session_store import save_session
|
||||
from backend.apps.agents.execution.prompt_builder import build_prompt_content
|
||||
from backend.apps.tools_lib.tools_lib import (
|
||||
_load_all as load_all_tools,
|
||||
load_builtin_permissions,
|
||||
)
|
||||
from backend.apps.analytics.collector import record as _analytics
|
||||
from backend.apps.agents.execution.agent_hooks import create_sdk_hooks
|
||||
|
||||
from claude_agent_sdk import (
|
||||
query, ClaudeAgentOptions, AssistantMessage, ResultMessage,
|
||||
query, ClaudeAgentOptions, AssistantMessage,
|
||||
)
|
||||
from claude_agent_sdk.types import (
|
||||
PermissionResultAllow, PermissionResultDeny,
|
||||
TextBlock, ToolUseBlock, StreamEvent, SystemMessage,
|
||||
)
|
||||
from backend.apps.agents.execution.agent_options import build_agent_options
|
||||
from claude_agent_sdk.types import StreamEvent
|
||||
from backend.apps.agents.manager.HaikFix.helpers.handle_stream_event import handle_stream_event
|
||||
from backend.apps.agents.manager.HaikFix.helpers.handle_assistant_message import handle_assistant_message
|
||||
|
||||
from backend.apps.agents.manager.HaikFix.PromptChunks import ImageChunk, ImageChunkDict, TextChunk, TextChunkDict
|
||||
from typing import List, Dict, Literal, Any, Union, Optional
|
||||
from typing import List, Dict, Literal, Union, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -47,7 +30,6 @@ def build_image_prompt_content(prompt: str, images: List[ImageChunk]) -> List[Te
|
||||
content.append(img.to_dict())
|
||||
return content
|
||||
|
||||
|
||||
PromptMsgDict = Dict[
|
||||
Literal["type", "message"],
|
||||
Dict[
|
||||
@@ -69,38 +51,45 @@ def build_prompt_msg(prompt: str, images: Optional[List[ImageChunk]]) -> PromptM
|
||||
}
|
||||
}
|
||||
|
||||
@typechecked
|
||||
async def run_agent_loop(
|
||||
prompt: str,
|
||||
images: list | None = None,
|
||||
images: Optional[List[ImageChunk]] = None,
|
||||
options: ClaudeAgentOptions | None = None,
|
||||
branch_id: str | None = None,
|
||||
):
|
||||
"""Run the Claude Agent SDK query loop for a session."""
|
||||
|
||||
prompt_msg = build_prompt_msg(prompt, images)
|
||||
prompt_msg = build_prompt_msg(prompt=prompt, images=images)
|
||||
|
||||
async def prompt_stream():
|
||||
yield prompt_msg
|
||||
|
||||
stream_text_msg_id = None
|
||||
stream_tool_msg_ids_ordered: list[str] = []
|
||||
stream_block_index_map: dict[int, str] = {}
|
||||
_turn_number = 0
|
||||
_first_event = True
|
||||
stream_text_msg_id: Optional[str] = None
|
||||
stream_tool_msg_ids_ordered: List[str] = []
|
||||
stream_block_index_map: Dict[int, str] = {}
|
||||
_turn_number: int = 0
|
||||
_first_event: bool = True
|
||||
|
||||
async for message in query(prompt=prompt_stream(), options=options):
|
||||
|
||||
if isinstance(message, StreamEvent):
|
||||
stream_text_msg_id = await _handle_stream_event(
|
||||
session_id, message.event,
|
||||
stream_text_msg_id, stream_tool_msg_ids_ordered, stream_block_index_map,
|
||||
stream_text_msg_id = await handle_stream_event(
|
||||
session_id=options.session_id,
|
||||
event=message.event,
|
||||
stream_text_msg_id=stream_text_msg_id,
|
||||
stream_tool_ids=stream_tool_msg_ids_ordered,
|
||||
block_map=stream_block_index_map,
|
||||
)
|
||||
|
||||
elif isinstance(message, AssistantMessage):
|
||||
stream_text_msg_id, stream_tool_msg_ids_ordered, stream_block_index_map = (
|
||||
await _handle_assistant_message(
|
||||
session, session_id, message, stream_text_msg_id,
|
||||
stream_tool_msg_ids_ordered, _turn_number,
|
||||
TextBlock, ToolUseBlock,
|
||||
await handle_assistant_message(
|
||||
session_id=options.session_id,
|
||||
branch_id=branch_id,
|
||||
message=message,
|
||||
stream_text_msg_id=stream_text_msg_id,
|
||||
stream_tool_ids=stream_tool_msg_ids_ordered,
|
||||
)
|
||||
)
|
||||
_turn_number += 1
|
||||
@@ -1,32 +1,38 @@
|
||||
from backend.apps.agents.manager.HaikFix.helpers.Message import Message, ToolCallContent
|
||||
from claude_agent_sdk.types import TextBlock, ToolUseBlock, AssistantMessage
|
||||
from backend.apps.agents.manager.ws_manager import ws_manager
|
||||
from claude_agent_sdk.types import (
|
||||
PermissionResultAllow, PermissionResultDeny,
|
||||
TextBlock, ToolUseBlock, StreamEvent, SystemMessage,
|
||||
)
|
||||
from typeguard import typechecked
|
||||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
||||
@typechecked
|
||||
async def handle_assistant_message(
|
||||
session, session_id, message, stream_text_msg_id,
|
||||
stream_tool_ids
|
||||
session_id: str,
|
||||
branch_id: str,
|
||||
message: AssistantMessage,
|
||||
stream_text_msg_id: str,
|
||||
stream_tool_ids: List[str],
|
||||
):
|
||||
content_parts = []
|
||||
tool_uses = []
|
||||
"""Handle an assistant message from the Claude Agent SDK."""
|
||||
# NOTE: Does not acctually save the message to the session.
|
||||
# TODO: Save the message to the session -Haik
|
||||
content_parts: List[str] = []
|
||||
tool_uses: List[ToolCallContent] = []
|
||||
for block in message.content:
|
||||
if isinstance(block, TextBlock):
|
||||
content_parts.append(block.text)
|
||||
elif isinstance(block, ToolUseBlock):
|
||||
tool_uses.append({"id": block.id, "tool": block.name, "input": block.input})
|
||||
tool_uses.append(ToolCallContent(id=block.id, tool=block.name, input=block.input))
|
||||
|
||||
if content_parts:
|
||||
asst_msg = Message(
|
||||
asst_msg: Message = Message(
|
||||
id=stream_text_msg_id or uuid4().hex,
|
||||
role="assistant", content="\n".join(content_parts),
|
||||
branch_id=session.active_branch_id,
|
||||
branch_id=branch_id,
|
||||
)
|
||||
session.messages.append(asst_msg)
|
||||
await ws_manager.emit_message(session_id, asst_msg)
|
||||
|
||||
for i, tu in enumerate(tool_uses):
|
||||
mid = stream_tool_ids[i] if i < len(stream_tool_ids) else uuid4().hex
|
||||
tool_msg = Message(id=mid, role="tool_call", content=tu, branch_id=session.active_branch_id)
|
||||
session.messages.append(tool_msg)
|
||||
for i, tu in enumerate[ToolCallContent](tool_uses):
|
||||
mid: str = stream_tool_ids[i] if i < len(stream_tool_ids) else uuid4().hex
|
||||
tool_msg: Message = Message(id=mid, role="tool_call", content=tu, branch_id=branch_id)
|
||||
await ws_manager.emit_message(session_id, tool_msg)
|
||||
@@ -0,0 +1,67 @@
|
||||
from backend.apps.agents.manager.ws_manager import ws_manager
|
||||
from typing import Any, Dict, Optional
|
||||
from typeguard import typechecked
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
@typechecked
|
||||
async def handle_stream_event(
|
||||
session_id: str,
|
||||
event: Dict[str, Any],
|
||||
stream_text_msg_id: Optional[str],
|
||||
stream_tool_ids: list[str],
|
||||
block_map: dict[int, str],
|
||||
) -> str | None:
|
||||
"""Process a single StreamEvent and return the (possibly updated) text msg id."""
|
||||
assert "type" in event, "Stream event missing 'type'"
|
||||
event_type: str = event["type"]
|
||||
|
||||
if event_type == "content_block_start":
|
||||
assert "index" in event, "content_block_start missing 'index'"
|
||||
assert "content_block" in event, "content_block_start missing 'content_block'"
|
||||
block: Dict[str, Any] = event["content_block"]
|
||||
index: int = event["index"]
|
||||
assert "type" in block, "content_block missing 'type'"
|
||||
block_type: str = block["type"]
|
||||
if block_type == "text":
|
||||
if stream_text_msg_id is None:
|
||||
stream_text_msg_id = uuid4().hex
|
||||
await ws_manager.emit_stream_start(session_id, stream_text_msg_id, "assistant")
|
||||
block_map[index] = stream_text_msg_id
|
||||
elif block_type == "tool_use":
|
||||
assert "name" in block, "tool_use content_block missing 'name'"
|
||||
tool_name: str = block["name"]
|
||||
tool_msg_id: str = uuid4().hex
|
||||
stream_tool_ids.append(tool_msg_id)
|
||||
block_map[index] = tool_msg_id
|
||||
await ws_manager.emit_stream_start(session_id, tool_msg_id, "tool_call", tool_name=tool_name)
|
||||
|
||||
elif event_type == "content_block_delta":
|
||||
assert "index" in event, "content_block_delta missing 'index'"
|
||||
assert "delta" in event, "content_block_delta missing 'delta'"
|
||||
index: int = event["index"]
|
||||
delta: Dict[str, Any] = event["delta"]
|
||||
msg_id: Optional[str] = block_map.get(index)
|
||||
if msg_id:
|
||||
assert "type" in delta, "delta missing 'type'"
|
||||
delta_type: str = delta["type"]
|
||||
if delta_type == "text_delta":
|
||||
assert "text" in delta, "text_delta missing 'text'"
|
||||
text: str = delta["text"]
|
||||
await ws_manager.emit_stream_delta(session_id, msg_id, text)
|
||||
elif delta_type == "input_json_delta":
|
||||
assert "partial_json" in delta, "input_json_delta missing 'partial_json'"
|
||||
partial_json: str = delta["partial_json"]
|
||||
await ws_manager.emit_stream_delta(session_id, msg_id, partial_json)
|
||||
|
||||
elif event_type == "content_block_stop":
|
||||
assert "index" in event, "content_block_stop missing 'index'"
|
||||
msg_id: Optional[str] = block_map.get(event["index"])
|
||||
if msg_id and msg_id != stream_text_msg_id:
|
||||
await ws_manager.emit_stream_end(session_id, msg_id)
|
||||
|
||||
elif event_type == "message_stop":
|
||||
if stream_text_msg_id:
|
||||
await ws_manager.emit_stream_end(session_id, stream_text_msg_id)
|
||||
|
||||
return stream_text_msg_id
|
||||
Reference in New Issue
Block a user