[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:
haikdc
2026-04-02 08:07:38 -07:00
parent b228b5cf72
commit ecdd215fe5
4 changed files with 159 additions and 141 deletions
+44 -88
View File
@@ -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