diff --git a/backend/apps/agents/manager/HaikFix/Agent.py b/backend/apps/agents/manager/HaikFix/Agent.py index 1507596b..9dd0e10a 100644 --- a/backend/apps/agents/manager/HaikFix/Agent.py +++ b/backend/apps/agents/manager/HaikFix/Agent.py @@ -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" diff --git a/backend/apps/agents/manager/HaikFix/agent_loop.py b/backend/apps/agents/manager/HaikFix/agent_loop.py index e7d4f90d..c331ced2 100644 --- a/backend/apps/agents/manager/HaikFix/agent_loop.py +++ b/backend/apps/agents/manager/HaikFix/agent_loop.py @@ -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 \ No newline at end of file diff --git a/backend/apps/agents/manager/HaikFix/helpers/handle_assistant_message.py b/backend/apps/agents/manager/HaikFix/helpers/handle_assistant_message.py index 5258570b..ecc41360 100644 --- a/backend/apps/agents/manager/HaikFix/helpers/handle_assistant_message.py +++ b/backend/apps/agents/manager/HaikFix/helpers/handle_assistant_message.py @@ -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) \ No newline at end of file diff --git a/backend/apps/agents/manager/HaikFix/helpers/handle_stream_event.py b/backend/apps/agents/manager/HaikFix/helpers/handle_stream_event.py index e69de29b..11e80d6c 100644 --- a/backend/apps/agents/manager/HaikFix/helpers/handle_stream_event.py +++ b/backend/apps/agents/manager/HaikFix/helpers/handle_stream_event.py @@ -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 \ No newline at end of file