From 781003227d2e433e7f3dba692a4f499edcc401e5 Mon Sep 17 00:00:00 2001 From: haikdc Date: Mon, 30 Mar 2026 02:42:48 -0700 Subject: [PATCH] [Haik]: Agentic refactor 3. DRY Up Cross-Cutting Patterns --- .cursor/mcp.json | 11 + backend/apps/agents/agent_hooks.py | 181 +++++ backend/apps/agents/agent_loop.py | 728 +++--------------- backend/apps/agents/agent_manager.py | 366 +-------- backend/apps/agents/agent_manager_meta.py | 126 +++ backend/apps/agents/agent_manager_ops.py | 232 ++++++ backend/apps/agents/agent_mock.py | 132 ++++ backend/apps/agents/agent_options.py | 217 ++++++ backend/apps/agents/agents.py | 166 ++-- backend/apps/agents/approval.py | 79 ++ backend/apps/agents/browser/executor.py | 28 +- backend/apps/agents/browser/runner.py | 31 +- .../apps/agents/browser_agent_mcp_server.py | 2 +- backend/apps/agents/browser_mcp_server.py | 2 +- .../apps/agents/invoke_agent_mcp_server.py | 2 +- backend/apps/agents/ws_manager.py | 71 ++ backend/apps/agents/ws_routes.py | 60 ++ backend/apps/subscriptions/__init__.py | 0 backend/apps/subscriptions/subscriptions.py | 192 +++++ backend/main.py | 186 +---- .../src/app/components/OnboardingModal.tsx | 10 +- frontend/src/app/pages/Settings/Settings.tsx | 12 +- 22 files changed, 1551 insertions(+), 1283 deletions(-) create mode 100644 .cursor/mcp.json create mode 100644 backend/apps/agents/agent_hooks.py create mode 100644 backend/apps/agents/agent_manager_meta.py create mode 100644 backend/apps/agents/agent_manager_ops.py create mode 100644 backend/apps/agents/agent_mock.py create mode 100644 backend/apps/agents/agent_options.py create mode 100644 backend/apps/agents/approval.py create mode 100644 backend/apps/agents/ws_routes.py create mode 100644 backend/apps/subscriptions/__init__.py create mode 100644 backend/apps/subscriptions/subscriptions.py diff --git a/.cursor/mcp.json b/.cursor/mcp.json new file mode 100644 index 00000000..77cfd6c6 --- /dev/null +++ b/.cursor/mcp.json @@ -0,0 +1,11 @@ +{ + "mcpServers": { + "mcp-docs-server": { + "command": "npx", + "args": [ + "-y", + "@assistant-ui/mcp-docs-server" + ] + } + } +} \ No newline at end of file diff --git a/backend/apps/agents/agent_hooks.py b/backend/apps/agents/agent_hooks.py new file mode 100644 index 00000000..0c5813e1 --- /dev/null +++ b/backend/apps/agents/agent_hooks.py @@ -0,0 +1,181 @@ +"""SDK hook factories for the agent loop. + +Creates the can_use_tool, pre_tool_hook, and post_tool_hook callables +required by ClaudeAgentOptions. +""" + +from __future__ import annotations + +import json +import logging +import re +import time +from datetime import datetime +from uuid import uuid4 + +from backend.apps.agents.models import AgentSession, Message +from backend.apps.agents.ws_manager import ws_manager +from backend.apps.agents.approval import request_approval +from backend.apps.agents.mcp_builder import get_effective_policy +from backend.apps.analytics.collector import record as _analytics + +logger = logging.getLogger(__name__) + + +def create_sdk_hooks( + session: AgentSession, + session_id: str, + sessions: dict[str, AgentSession], + builtin_perms: dict, + PermissionResultAllow, + PermissionResultDeny, +): + """Return (can_use_tool, pre_tool_hook, post_tool_hook) closures.""" + + tool_start_times: dict[str, float] = {} + + async def _request_user_approval(tool_name: str, tool_input) -> dict: + safe_input = tool_input if isinstance(tool_input, dict) else {} + return await request_approval(session, tool_name, safe_input, track_analytics=True) + + async def can_use_tool(tool_name, input_data, context): + if tool_name != "AskUserQuestion": + policy = get_effective_policy(tool_name, builtin_perms) + if policy == "always_allow": + return PermissionResultAllow(updated_input=input_data) + if policy == "deny": + return PermissionResultDeny(message="Tool denied by permission policy") + decision = await _request_user_approval(tool_name, input_data) + if decision.get("behavior") == "allow": + return PermissionResultAllow(updated_input=decision.get("updated_input", input_data)) + return PermissionResultDeny(message=decision.get("message", "User denied this action")) + + async def pre_tool_hook(input_data, tool_use_id, context): + tool_name = input_data.get("tool_name", "") + hook_event = input_data.get("hook_event_name", "PreToolUse") + if tool_name and tool_name != "AskUserQuestion": + policy = get_effective_policy(tool_name, builtin_perms) + if policy == "deny": + return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "deny", "permissionDecisionReason": "Tool denied by permission policy"}} + if policy == "ask": + tool_input = input_data.get("tool_input", {}) + decision = await _request_user_approval(tool_name, tool_input) + if decision.get("behavior") == "allow": + if tool_use_id: + tool_start_times[tool_use_id] = time.time() + return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "allow"}} + return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "deny", "permissionDecisionReason": decision.get("message", "User denied this action")}} + if tool_use_id: + tool_start_times[tool_use_id] = time.time() + return {} + + async def post_tool_hook(input_data, tool_use_id, context): + elapsed_ms = None + if tool_use_id and tool_use_id in tool_start_times: + elapsed_ms = int((time.time() - tool_start_times.pop(tool_use_id)) * 1000) + + raw_response = input_data.get("tool_response", "") + + hook_tool_name_early = input_data.get("tool_name", "") + if hook_tool_name_early: + _is_mcp = "__" in hook_tool_name_early + _mcp_server = "" + _tool_short = hook_tool_name_early + if _is_mcp: + _mcp_match = re.match(r"mcp__([^_]+(?:-[^_]+)*)__(.+)", hook_tool_name_early) + if _mcp_match: + _mcp_server = _mcp_match.group(1) + _tool_short = _mcp_match.group(2) + _analytics("tool.executed", { + "tool_name": hook_tool_name_early, "tool_short_name": _tool_short, + "tool_type": "mcp" if _is_mcp else "builtin", "mcp_server": _mcp_server, + "duration_ms": elapsed_ms, + "success": not (isinstance(raw_response, str) and raw_response.startswith("Error")), + "model": session.model, "provider": session.provider, + }, session_id=session_id, dashboard_id=session.dashboard_id) + + if isinstance(raw_response, list) and raw_response: + text_parts = [b.get("text", "") for b in raw_response if isinstance(b, dict) and b.get("type") == "text"] + if text_parts: + raw_response = "\n".join(text_parts) if len(text_parts) > 1 else text_parts[0] + + if isinstance(raw_response, str): + content = raw_response + else: + try: + content = json.dumps(raw_response, indent=2, default=str) + except Exception: + content = str(raw_response) + + result_payload: dict = {"text": content} + hook_tool_name = input_data.get("tool_name", "") + if hook_tool_name: + result_payload["tool_name"] = hook_tool_name + if elapsed_ms is not None: + result_payload["elapsed_ms"] = elapsed_ms + + if hook_tool_name == "Agent": + sub_payload = _build_sub_agent_session( + input_data, raw_response, content, session, session_id, sessions, + ) + if sub_payload: + result_payload["sub_session_id"] = sub_payload + + result_msg = Message(role="tool_result", content=result_payload, branch_id=session.active_branch_id) + session.messages.append(result_msg) + await ws_manager.emit_message(session_id, result_msg) + return {"continue_": True} + + return can_use_tool, pre_tool_hook, post_tool_hook + + +async def _broadcast_sub_session(sub_session: AgentSession): + await ws_manager.broadcast_global("agent:status", { + "session_id": sub_session.id, "status": sub_session.status, + "session": sub_session.model_dump(mode="json"), + }) + + +def _build_sub_agent_session( + input_data: dict, raw_response, content: str, + session: AgentSession, session_id: str, + sessions: dict[str, AgentSession], +) -> str | None: + """Create a sub-agent session from an Agent tool result. Returns sub_session_id or None.""" + tool_input = input_data.get("tool_input", {}) + agent_prompt = tool_input.get("prompt", tool_input.get("task", "")) + sub_text = content + sub_cost = 0.0 + sub_tokens: dict = {"input": 0, "output": 0} + sub_model = session.model + if isinstance(raw_response, dict): + blocks = raw_response.get("content") + if isinstance(blocks, list): + parts = [b.get("text", "") for b in blocks if isinstance(b, dict) and b.get("type") == "text"] + if parts: + sub_text = "\n".join(parts) if len(parts) > 1 else parts[0] + elif isinstance(raw_response.get("text"), str): + sub_text = raw_response["text"] + usage = raw_response.get("usage", {}) + if isinstance(usage, dict): + sub_tokens["input"] = usage.get("input_tokens", 0) + usage.get("cache_creation_input_tokens", 0) + usage.get("cache_read_input_tokens", 0) + sub_tokens["output"] = usage.get("output_tokens", 0) + if raw_response.get("model"): + sub_model = raw_response["model"] + + sub_session_id = uuid4().hex + sub_name = agent_prompt[:50] if agent_prompt else "Sub-agent" + sub_session = AgentSession( + id=sub_session_id, name=sub_name, status="completed", model=sub_model, + mode="sub-agent", cwd=session.cwd, created_at=datetime.now(), + cost_usd=sub_cost, tokens=sub_tokens, + messages=[ + Message(role="user", content=agent_prompt, branch_id="main"), + Message(role="assistant", content=sub_text, branch_id="main"), + ], + dashboard_id=session.dashboard_id, parent_session_id=session_id, + ) + sessions[sub_session_id] = sub_session + import asyncio + asyncio.ensure_future(_broadcast_sub_session(sub_session)) + return sub_session_id diff --git a/backend/apps/agents/agent_loop.py b/backend/apps/agents/agent_loop.py index 007216c4..4b0b82e7 100644 --- a/backend/apps/agents/agent_loop.py +++ b/backend/apps/agents/agent_loop.py @@ -1,200 +1,31 @@ -"""Main agent loop — extracted from AgentManager._run_agent_loop. +"""Main agent loop — orchestrates the Claude Agent SDK query loop. -Handles the Claude Agent SDK query loop, approval hooks, streaming, -mock-agent fallback, and session-completed analytics. +Heavy logic is delegated to sibling modules: +- agent_mock – mock-agent fallback, streaming helpers, session analytics +- agent_hooks – SDK hook factories (approval, permissions, post-tool) +- agent_options – MCP server construction & ClaudeAgentOptions building """ from __future__ import annotations import asyncio -import json import logging -import os -import sys -import time -from datetime import datetime from uuid import uuid4 -from backend.apps.agents.models import AgentSession, ApprovalRequest, Message +from backend.apps.agents.models import AgentSession, Message from backend.apps.agents.ws_manager import ws_manager -from backend.apps.agents.prompt_builder import ( - resolve_mode, compose_system_prompt, build_connected_tools_context, - build_outputs_context, build_browser_context, build_prompt_content, - get_pre_selected_browser_ids, -) -from backend.apps.agents.mcp_builder import ( - FULL_TOOLS, build_mcp_servers, get_effective_policy, get_all_tool_names, - _get_denied_tool_names, _get_all_known_tool_names, _is_fully_denied, -) from backend.apps.agents.session_store import save_session -from backend.apps.settings.settings import load_settings +from backend.apps.agents.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.common.mcp_utils import sanitize_server_name as _sanitize_server_name from backend.apps.analytics.collector import record as _analytics +from backend.apps.agents.agent_mock import run_mock_agent, fire_session_completed logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Streaming helpers -# --------------------------------------------------------------------------- - -async def stream_text(session_id: str, msg_id: str, text: str, delay: float = 0.03): - await ws_manager.send_to_session(session_id, "agent:stream_start", { - "session_id": session_id, "message_id": msg_id, "role": "assistant", - }) - words = text.split(" ") - for i, word in enumerate(words): - chunk = word if i == 0 else " " + word - await ws_manager.send_to_session(session_id, "agent:stream_delta", { - "session_id": session_id, "message_id": msg_id, "delta": chunk, - }) - await asyncio.sleep(delay) - await ws_manager.send_to_session(session_id, "agent:stream_end", { - "session_id": session_id, "message_id": msg_id, - }) - - -async def stream_tool_input(session_id: str, msg_id: str, tool_name: str, input_json: str, delay: float = 0.02): - await ws_manager.send_to_session(session_id, "agent:stream_start", { - "session_id": session_id, "message_id": msg_id, "role": "tool_call", "tool_name": tool_name, - }) - chunk_size = 12 - for i in range(0, len(input_json), chunk_size): - await ws_manager.send_to_session(session_id, "agent:stream_delta", { - "session_id": session_id, "message_id": msg_id, "delta": input_json[i:i + chunk_size], - }) - await asyncio.sleep(delay) - await ws_manager.send_to_session(session_id, "agent:stream_end", { - "session_id": session_id, "message_id": msg_id, - }) - - -# --------------------------------------------------------------------------- -# Analytics helper -# --------------------------------------------------------------------------- - -def fire_session_completed(session: AgentSession, sessions_dict: dict[str, AgentSession]): - duration = 0.0 - if session.created_at: - end = session.closed_at or datetime.now() - duration = (end - session.created_at).total_seconds() - tool_names = [ - m.content.get("tool", "") for m in session.messages - if m.role == "tool_call" and isinstance(m.content, dict) - ] - user_messages = [ - (m.content if isinstance(m.content, str) else str(m.content))[:200] - for m in session.messages if m.role == "user" - ] - _analytics("session.completed", { - "model": session.model, - "provider": getattr(session, "provider", "anthropic"), - "mode": session.mode, - "cost_usd": session.cost_usd, - "message_count": len([m for m in session.messages if m.role in ("user", "assistant")]), - "duration_seconds": round(duration, 1), - "status": session.status, - "tool_count": len(tool_names), - "tools_list": list(set(tool_names)), - "session_title": session.name, - "first_user_message": user_messages[0] if user_messages else "", - "input_tokens": session.tokens.get("input", 0), - "output_tokens": session.tokens.get("output", 0), - "is_sub_agent": session.parent_session_id is not None, - "parent_session_id": session.parent_session_id, - "sub_agent_count": len([s for s in sessions_dict.values() if s.parent_session_id == session.id]), - "branch_count": len(session.branches), - }, session_id=session.id, dashboard_id=session.dashboard_id) - - -# --------------------------------------------------------------------------- -# Mock agent -# --------------------------------------------------------------------------- - -async def run_mock_agent(session_id: str, prompt: str, sessions: dict[str, AgentSession]): - session = sessions.get(session_id) - if not session: - return - - await asyncio.sleep(1) - - request_id = uuid4().hex - approval_req = ApprovalRequest( - id=request_id, session_id=session_id, tool_name="Bash", - tool_input={"command": f"echo 'Processing: {prompt}'", "description": "Echo the user prompt"}, - ) - session.pending_approvals.append(approval_req) - session.status = "waiting_approval" - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "waiting_approval", - }) - - decision = await ws_manager.send_approval_request( - session_id, request_id, "Bash", - {"command": f"echo 'Processing: {prompt}'", "description": "Echo the user prompt"}, - ) - - session.pending_approvals = [a for a in session.pending_approvals if a.id != request_id] - session.status = "running" - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "running", - }) - - tool_input_content = {"tool": "Bash", "input": {"command": f"echo 'Processing: {prompt}'"}, "approved": decision.get("behavior") == "allow"} - tool_msg_id = uuid4().hex - await stream_tool_input(session_id, tool_msg_id, "Bash", json.dumps(tool_input_content["input"], indent=2)) - tool_msg = Message(id=tool_msg_id, role="tool_call", content=tool_input_content, branch_id=session.active_branch_id) - session.messages.append(tool_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": tool_msg.model_dump(mode="json"), - }) - - await asyncio.sleep(1) - - if decision.get("behavior") == "allow": - tool_result = Message(role="tool_result", content=f"Processing: {prompt}", branch_id=session.active_branch_id) - session.messages.append(tool_result) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": tool_result.model_dump(mode="json"), - }) - - await asyncio.sleep(1) - - asst_text = ( - f"I've processed your request: \"{prompt}\"\n\n" - "This is a mock response because `claude-agent-sdk` is not installed. " - "Install it with `pip install claude-agent-sdk` to use real Claude Code instances.\n\n" - f"The agent was configured with:\n- Model: {session.model}\n- Mode: {session.mode}" - ) - asst_msg_id = uuid4().hex - await stream_text(session_id, asst_msg_id, asst_text) - - asst_msg = Message(id=asst_msg_id, role="assistant", content=asst_text, branch_id=session.active_branch_id) - session.messages.append(asst_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": asst_msg.model_dump(mode="json"), - }) - - session.status = "completed" - session.closed_at = datetime.now() - session.cost_usd = 0.001 - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "completed", - "session": session.model_dump(mode="json"), - }) - await ws_manager.send_to_session(session_id, "agent:cost_update", { - "session_id": session_id, "cost_usd": session.cost_usd, - }) - - -# --------------------------------------------------------------------------- -# Main agent loop -# --------------------------------------------------------------------------- - async def run_agent_loop( sessions: dict[str, AgentSession], session_id: str, @@ -222,8 +53,7 @@ async def run_agent_loop( ) from claude_agent_sdk.types import ( HookMatcher, PermissionResultAllow, PermissionResultDeny, - TextBlock, ToolUseBlock, StreamEvent, - SystemMessage, + TextBlock, ToolUseBlock, StreamEvent, SystemMessage, ) except ImportError: logger.warning("claude_agent_sdk not installed, running in mock mode") @@ -231,327 +61,21 @@ async def run_agent_loop( return session.status = "running" - _builtin_perms = load_builtin_permissions() + builtin_perms = load_builtin_permissions() - async def _request_user_approval(tool_name: str, tool_input) -> dict: - safe_input = tool_input if isinstance(tool_input, dict) else {} - request_id = uuid4().hex - approval_req = ApprovalRequest( - id=request_id, session_id=session_id, tool_name=tool_name, tool_input=safe_input, - ) - session.pending_approvals.append(approval_req) - session.status = "waiting_approval" + from backend.apps.agents.agent_hooks import create_sdk_hooks + can_use_tool, pre_tool_hook, post_tool_hook = create_sdk_hooks( + session, session_id, sessions, builtin_perms, + PermissionResultAllow, PermissionResultDeny, + ) - _analytics("approval.requested", { - "tool_name": tool_name, - "is_first_approval_in_session": len(session.pending_approvals) == 1, - "model": session.model, - }, session_id=session_id, dashboard_id=session.dashboard_id) - - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "waiting_approval", - }) - - decision = await ws_manager.send_approval_request( - session_id, request_id, tool_name, safe_input, - ) - - approval_latency_ms = int((datetime.now() - approval_req.created_at).total_seconds() * 1000) - _analytics("approval.resolved", { - "tool_name": tool_name, - "decision": decision.get("behavior", "unknown"), - "latency_ms": approval_latency_ms, - "input_was_modified": decision.get("updated_input") is not None, - "model": session.model, - }, session_id=session_id, dashboard_id=session.dashboard_id) - - session.pending_approvals = [a for a in session.pending_approvals if a.id != request_id] - session.status = "running" - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "running", - }) - return decision - - async def can_use_tool(tool_name, input_data, context): - if tool_name != "AskUserQuestion": - policy = get_effective_policy(tool_name, _builtin_perms) - if policy == "always_allow": - return PermissionResultAllow(updated_input=input_data) - if policy == "deny": - return PermissionResultDeny(message="Tool denied by permission policy") - decision = await _request_user_approval(tool_name, input_data) - if decision.get("behavior") == "allow": - return PermissionResultAllow(updated_input=decision.get("updated_input", input_data)) - return PermissionResultDeny(message=decision.get("message", "User denied this action")) - - tool_start_times: dict[str, float] = {} - - async def pre_tool_hook(input_data, tool_use_id, context): - tool_name = input_data.get("tool_name", "") - hook_event = input_data.get("hook_event_name", "PreToolUse") - if tool_name and tool_name != "AskUserQuestion": - policy = get_effective_policy(tool_name, _builtin_perms) - if policy == "deny": - return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "deny", "permissionDecisionReason": "Tool denied by permission policy"}} - if policy == "ask": - tool_input = input_data.get("tool_input", {}) - decision = await _request_user_approval(tool_name, tool_input) - if decision.get("behavior") == "allow": - if tool_use_id: - tool_start_times[tool_use_id] = time.time() - return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "allow"}} - return {"hookSpecificOutput": {"hookEventName": hook_event, "permissionDecision": "deny", "permissionDecisionReason": decision.get("message", "User denied this action")}} - if tool_use_id: - tool_start_times[tool_use_id] = time.time() - return {} - - async def post_tool_hook(input_data, tool_use_id, context): - import re as _re_tool - elapsed_ms = None - if tool_use_id and tool_use_id in tool_start_times: - elapsed_ms = int((time.time() - tool_start_times.pop(tool_use_id)) * 1000) - - raw_response = input_data.get("tool_response", "") - - hook_tool_name_early = input_data.get("tool_name", "") - if hook_tool_name_early: - _is_mcp = "__" in hook_tool_name_early - _mcp_server = "" - _tool_short = hook_tool_name_early - if _is_mcp: - _mcp_match = _re_tool.match(r"mcp__([^_]+(?:-[^_]+)*)__(.+)", hook_tool_name_early) - if _mcp_match: - _mcp_server = _mcp_match.group(1) - _tool_short = _mcp_match.group(2) - _analytics("tool.executed", { - "tool_name": hook_tool_name_early, "tool_short_name": _tool_short, - "tool_type": "mcp" if _is_mcp else "builtin", "mcp_server": _mcp_server, - "duration_ms": elapsed_ms, - "success": not (isinstance(raw_response, str) and raw_response.startswith("Error")), - "model": session.model, "provider": session.provider, - }, session_id=session_id, dashboard_id=session.dashboard_id) - - if isinstance(raw_response, list) and raw_response: - text_parts = [b.get("text", "") for b in raw_response if isinstance(b, dict) and b.get("type") == "text"] - if text_parts: - raw_response = "\n".join(text_parts) if len(text_parts) > 1 else text_parts[0] - - if isinstance(raw_response, str): - content = raw_response - else: - try: - content = json.dumps(raw_response, indent=2, default=str) - except Exception: - content = str(raw_response) - - result_payload: dict = {"text": content} - hook_tool_name = input_data.get("tool_name", "") - if hook_tool_name: - result_payload["tool_name"] = hook_tool_name - if elapsed_ms is not None: - result_payload["elapsed_ms"] = elapsed_ms - - if hook_tool_name == "Agent": - tool_input = input_data.get("tool_input", {}) - agent_prompt = tool_input.get("prompt", tool_input.get("task", "")) - sub_text = content - sub_cost = 0.0 - sub_tokens: dict = {"input": 0, "output": 0} - sub_model = session.model - if isinstance(raw_response, dict): - blocks = raw_response.get("content") - if isinstance(blocks, list): - parts = [b.get("text", "") for b in blocks if isinstance(b, dict) and b.get("type") == "text"] - if parts: - sub_text = "\n".join(parts) if len(parts) > 1 else parts[0] - elif isinstance(raw_response.get("text"), str): - sub_text = raw_response["text"] - usage = raw_response.get("usage", {}) - if isinstance(usage, dict): - sub_tokens["input"] = usage.get("input_tokens", 0) + usage.get("cache_creation_input_tokens", 0) + usage.get("cache_read_input_tokens", 0) - sub_tokens["output"] = usage.get("output_tokens", 0) - if raw_response.get("model"): - sub_model = raw_response["model"] - - sub_session_id = uuid4().hex - sub_name = agent_prompt[:50] if agent_prompt else "Sub-agent" - sub_session = AgentSession( - id=sub_session_id, name=sub_name, status="completed", model=sub_model, - mode="sub-agent", cwd=session.cwd, created_at=datetime.now(), - cost_usd=sub_cost, tokens=sub_tokens, - messages=[ - Message(role="user", content=agent_prompt, branch_id="main"), - Message(role="assistant", content=sub_text, branch_id="main"), - ], - dashboard_id=session.dashboard_id, parent_session_id=session_id, - ) - sessions[sub_session_id] = sub_session - await ws_manager.broadcast_global("agent:status", { - "session_id": sub_session_id, "status": sub_session.status, - "session": sub_session.model_dump(mode="json"), - }) - result_payload["sub_session_id"] = sub_session_id - - result_msg = Message(role="tool_result", content=result_payload, branch_id=session.active_branch_id) - session.messages.append(result_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": result_msg.model_dump(mode="json"), - }) - return {"continue_": True} + from backend.apps.agents.agent_options import build_agent_options try: - _, mode_sys_prompt, _ = resolve_mode(session.mode, get_all_tool_names) - connected_tools_ctx = build_connected_tools_context( - session.allowed_tools, load_all_tools, get_all_tool_names, _is_fully_denied, _get_denied_tool_names, + options_kwargs = await build_agent_options( + session, builtin_perms, can_use_tool, pre_tool_hook, post_tool_hook, + fork_session=fork_session, selected_browser_ids=selected_browser_ids, ) - outputs_ctx = build_outputs_context() - browser_ctx = build_browser_context(session.dashboard_id, selected_browser_ids=selected_browser_ids) - global_settings = load_settings() - composed_prompt = compose_system_prompt( - global_settings.default_system_prompt, mode_sys_prompt, session.system_prompt, - connected_tools_ctx, outputs_ctx, browser_ctx, - ) - - if session.mode == "view-builder": - from backend.apps.outputs.view_builder_templates import VIEW_BUILDER_SKILL - skill_block = f"\n{VIEW_BUILDER_SKILL}\n" - composed_prompt = f"{composed_prompt}\n\n{skill_block}" if composed_prompt else skill_block - - mcp_servers = await build_mcp_servers(session.allowed_tools) - - _browser_delegation_tools = ["CreateBrowserAgent", "BrowserAgent", "BrowserAgents"] - _browser_all_denied = all(_builtin_perms.get(t, "always_allow") == "deny" for t in _browser_delegation_tools) - - if not _browser_all_denied: - browser_agent_server_path = os.path.join(os.path.dirname(__file__), "browser_agent_mcp_server.py") - backend_port = os.environ.get("OPENSWARM_PORT", "8324") - pre_selected_bids = get_pre_selected_browser_ids(session.dashboard_id) - mcp_servers["openswarm-browser-agent"] = { - "command": sys.executable, - "args": [browser_agent_server_path], - "env": { - "OPENSWARM_PORT": backend_port, - "OPENSWARM_AGENT_MODEL": session.model, - "OPENSWARM_DASHBOARD_ID": session.dashboard_id or "", - "OPENSWARM_PRE_SELECTED_BROWSER_IDS": ",".join(pre_selected_bids), - "OPENSWARM_PARENT_SESSION_ID": session.id, - }, - "type": "stdio", - } - - _invoke_agent_tools = ["InvokeAgent"] - _invoke_all_denied = all(_builtin_perms.get(t, "always_allow") == "deny" for t in _invoke_agent_tools) - - if not _invoke_all_denied: - invoke_agent_server_path = os.path.join(os.path.dirname(__file__), "invoke_agent_mcp_server.py") - backend_port = os.environ.get("OPENSWARM_PORT", "8324") - mcp_servers["openswarm-invoke-agent"] = { - "command": sys.executable, - "args": [invoke_agent_server_path], - "env": { - "OPENSWARM_PORT": backend_port, - "OPENSWARM_PARENT_SESSION_ID": session.id, - "OPENSWARM_DASHBOARD_ID": session.dashboard_id or "", - }, - "type": "stdio", - } - - effective_allowed = [ - t for t in session.allowed_tools - if t in FULL_TOOLS and _builtin_perms.get(t, "always_allow") == "always_allow" - ] - effective_disallowed = [ - t for t in FULL_TOOLS - if _builtin_perms.get(t, "always_allow") == "deny" - ] - - if mcp_servers: - all_tools_list = load_all_tools() - for name in mcp_servers: - if name == "openswarm-browser-agent": - for bt in _browser_delegation_tools: - policy = _builtin_perms.get(bt, "always_allow") - if policy == "always_allow": - effective_allowed.append(f"mcp__openswarm-browser-agent__{bt}") - elif policy == "deny": - effective_disallowed.append(f"mcp__openswarm-browser-agent__{bt}") - continue - if name == "openswarm-invoke-agent": - for it in _invoke_agent_tools: - policy = _builtin_perms.get(it, "always_allow") - if policy == "always_allow": - effective_allowed.append(f"mcp__openswarm-invoke-agent__{it}") - elif policy == "deny": - effective_disallowed.append(f"mcp__openswarm-invoke-agent__{it}") - continue - tool_def = next( - (t for t in all_tools_list if t.mcp_config and t.enabled and _sanitize_server_name(t.name) == name), - None, - ) - if tool_def: - denied = _get_denied_tool_names(tool_def) - known = _get_all_known_tool_names(tool_def) - for tn in known - denied: - policy = tool_def.tool_permissions.get(tn, "ask") - if policy == "always_allow": - effective_allowed.append(f"mcp__{name}__{tn}") - for tn in denied: - effective_disallowed.append(f"mcp__{name}__{tn}") - else: - effective_allowed.append(f"mcp__{name}__*") - - google_allowed = [t for t in effective_allowed if "google-workspace" in t] - reddit_allowed = [t for t in effective_allowed if "reddit" in t] - builtin_allowed = [t for t in effective_allowed if not t.startswith("mcp__")] - logger.info(f"[MCP-DEBUG] effective_allowed: {len(effective_allowed)} total " - f"(builtins={len(builtin_allowed)}, google={len(google_allowed)}, reddit={len(reddit_allowed)})") - if effective_disallowed: - logger.info(f"[MCP-DEBUG] effective_disallowed: {effective_disallowed}") - - options_kwargs: dict = { - "model": session.model, - "max_buffer_size": 5 * 1024 * 1024, - "permission_mode": "default", - "can_use_tool": can_use_tool, - "hooks": { - "PreToolUse": [HookMatcher(matcher=None, hooks=[pre_tool_hook])], - "PostToolUse": [HookMatcher(matcher=None, hooks=[post_tool_hook])], - }, - "allowed_tools": effective_allowed, - "disallowed_tools": effective_disallowed, - "include_partial_messages": True, - } - - from backend.apps.nine_router import is_running as _9r_running - if global_settings.anthropic_api_key: - options_kwargs["env"] = {"ANTHROPIC_API_KEY": global_settings.anthropic_api_key} - logger.info("[MCP-DEBUG] Using direct API key") - elif _9r_running(): - options_kwargs["env"] = { - "ANTHROPIC_API_KEY": "9router", - "ANTHROPIC_BASE_URL": "http://localhost:20128", - } - options_kwargs["extra_args"] = {"bare": None} - logger.info("[MCP-DEBUG] Using 9Router (bare mode)") - else: - raise ValueError("No AI provider configured. Set an API key or connect a subscription.") - - if mcp_servers: - options_kwargs["mcp_servers"] = mcp_servers - mcp_json_len = len(json.dumps({"mcpServers": mcp_servers})) - logger.info(f"[MCP-DEBUG] mcp_servers passed to SDK: {list(mcp_servers.keys())}, JSON length={mcp_json_len}") - if composed_prompt: - options_kwargs["system_prompt"] = composed_prompt - if session.max_turns: - options_kwargs["max_turns"] = session.max_turns - if session.cwd: - options_kwargs["cwd"] = session.cwd - if session.sdk_session_id: - options_kwargs["resume"] = session.sdk_session_id - if fork_session: - options_kwargs["fork_session"] = True - - logger.info(f"[MCP-DEBUG] Creating ClaudeAgentOptions with model={session.model}") options = ClaudeAgentOptions(**options_kwargs) logger.info("[MCP-DEBUG] ClaudeAgentOptions created. Starting query...") @@ -574,110 +98,23 @@ async def run_agent_loop( logger.info(f"[MCP-DEBUG] SystemMessage: {raw}") if isinstance(message, StreamEvent): - event = message.event - event_type = event.get("type") - - if event_type == "content_block_start": - block = event.get("content_block", {}) - index = event.get("index") - block_type = block.get("type") - if block_type == "text": - if stream_text_msg_id is None: - stream_text_msg_id = uuid4().hex - await ws_manager.send_to_session(session_id, "agent:stream_start", { - "session_id": session_id, "message_id": stream_text_msg_id, "role": "assistant", - }) - stream_block_index_map[index] = stream_text_msg_id - elif block_type == "tool_use": - tool_msg_id = uuid4().hex - stream_tool_msg_ids_ordered.append(tool_msg_id) - stream_block_index_map[index] = tool_msg_id - await ws_manager.send_to_session(session_id, "agent:stream_start", { - "session_id": session_id, "message_id": tool_msg_id, - "role": "tool_call", "tool_name": block.get("name", ""), - }) - - elif event_type == "content_block_delta": - index = event.get("index") - delta = event.get("delta", {}) - delta_type = delta.get("type") - msg_id = stream_block_index_map.get(index) - if msg_id and delta_type == "text_delta": - await ws_manager.send_to_session(session_id, "agent:stream_delta", { - "session_id": session_id, "message_id": msg_id, "delta": delta.get("text", ""), - }) - elif msg_id and delta_type == "input_json_delta": - await ws_manager.send_to_session(session_id, "agent:stream_delta", { - "session_id": session_id, "message_id": msg_id, "delta": delta.get("partial_json", ""), - }) - - elif event_type == "content_block_stop": - index = event.get("index") - msg_id = stream_block_index_map.get(index) - if msg_id and msg_id != stream_text_msg_id: - await ws_manager.send_to_session(session_id, "agent:stream_end", { - "session_id": session_id, "message_id": msg_id, - }) - - elif event_type == "message_stop": - if stream_text_msg_id: - await ws_manager.send_to_session(session_id, "agent:stream_end", { - "session_id": session_id, "message_id": stream_text_msg_id, - }) + 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, + ) elif isinstance(message, AssistantMessage): - content_parts = [] - tool_uses = [] - 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}) - - if content_parts: - asst_msg = Message( - id=stream_text_msg_id or uuid4().hex, - role="assistant", content="\n".join(content_parts), - branch_id=session.active_branch_id, + 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, ) - session.messages.append(asst_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": asst_msg.model_dump(mode="json"), - }) - - for i, tu in enumerate(tool_uses): - mid = stream_tool_msg_ids_ordered[i] if i < len(stream_tool_msg_ids_ordered) else uuid4().hex - tool_msg = Message(id=mid, role="tool_call", content=tu, branch_id=session.active_branch_id) - session.messages.append(tool_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": tool_msg.model_dump(mode="json"), - }) - + ) _turn_number += 1 - _analytics("turn.completed", { - "turn_number": _turn_number, "tool_calls_in_turn": len(tool_uses), "model": session.model, - }, session_id=session_id, dashboard_id=session.dashboard_id) - - stream_text_msg_id = None - stream_tool_msg_ids_ordered = [] - stream_block_index_map = {} elif isinstance(message, ResultMessage): - session.sdk_session_id = getattr(message, "session_id", None) - cost = getattr(message, "total_cost_usd", None) - if cost is not None: - session.cost_usd = cost - await ws_manager.send_to_session(session_id, "agent:cost_update", { - "session_id": session_id, "cost_usd": session.cost_usd, - }) - usage = getattr(message, "usage", None) or {} - if isinstance(usage, dict): - inp = usage.get("input_tokens", 0) or 0 - out = usage.get("output_tokens", 0) or 0 - cache_create = usage.get("cache_creation_input_tokens", 0) or 0 - cache_read = usage.get("cache_read_input_tokens", 0) or 0 - session.tokens["input"] = inp + cache_create + cache_read - session.tokens["output"] = out + await _handle_result_message(session, session_id, message) session.status = "completed" except asyncio.CancelledError: @@ -691,16 +128,107 @@ async def run_agent_loop( }, session_id=session_id, dashboard_id=session.dashboard_id) error_msg = Message(role="system", content=f"Error: {str(e)}", branch_id=session.active_branch_id) session.messages.append(error_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": error_msg.model_dump(mode="json"), - }) + await ws_manager.emit_message(session_id, error_msg) finally: if session_id in sessions: - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": session.status, - "session": session.model_dump(mode="json"), - }) + await ws_manager.emit_status(session_id, session.status, session) try: save_session(session_id, session.model_dump(mode="json")) except Exception as e: logger.warning(f"Failed to snapshot session {session_id}: {e}") + + +async def _handle_stream_event( + session_id: str, event: dict, + stream_text_msg_id: str | None, + stream_tool_ids: list[str], + block_map: dict[int, str], +) -> str | None: + """Process a single StreamEvent and return the (possibly updated) text msg id.""" + event_type = event.get("type") + + if event_type == "content_block_start": + block = event.get("content_block", {}) + index = event.get("index") + if block.get("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.get("type") == "tool_use": + tool_msg_id = 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=block.get("name", "")) + + elif event_type == "content_block_delta": + index = event.get("index") + delta = event.get("delta", {}) + msg_id = block_map.get(index) + if msg_id: + delta_type = delta.get("type") + if delta_type == "text_delta": + await ws_manager.emit_stream_delta(session_id, msg_id, delta.get("text", "")) + elif delta_type == "input_json_delta": + await ws_manager.emit_stream_delta(session_id, msg_id, delta.get("partial_json", "")) + + elif event_type == "content_block_stop": + msg_id = block_map.get(event.get("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 + + +async def _handle_assistant_message( + session, session_id, message, stream_text_msg_id, + stream_tool_ids, turn_number, TextBlock, ToolUseBlock, +): + content_parts = [] + tool_uses = [] + 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}) + + if content_parts: + asst_msg = Message( + id=stream_text_msg_id or uuid4().hex, + role="assistant", content="\n".join(content_parts), + branch_id=session.active_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) + await ws_manager.emit_message(session_id, tool_msg) + + _analytics("turn.completed", { + "turn_number": turn_number + 1, "tool_calls_in_turn": len(tool_uses), "model": session.model, + }, session_id=session_id, dashboard_id=session.dashboard_id) + + return None, [], {} + + +async def _handle_result_message(session, session_id, message): + session.sdk_session_id = getattr(message, "session_id", None) + cost = getattr(message, "total_cost_usd", None) + if cost is not None: + session.cost_usd = cost + await ws_manager.emit_cost_update(session_id, session.cost_usd) + usage = getattr(message, "usage", None) or {} + if isinstance(usage, dict): + inp = usage.get("input_tokens", 0) or 0 + out = usage.get("output_tokens", 0) or 0 + cache_create = usage.get("cache_creation_input_tokens", 0) or 0 + cache_read = usage.get("cache_read_input_tokens", 0) or 0 + session.tokens["input"] = inp + cache_create + cache_read + session.tokens["output"] = out diff --git a/backend/apps/agents/agent_manager.py b/backend/apps/agents/agent_manager.py index aca46b49..b0a2ce7d 100644 --- a/backend/apps/agents/agent_manager.py +++ b/backend/apps/agents/agent_manager.py @@ -1,10 +1,12 @@ """Thin coordinator for agent sessions. Heavy logic lives in sibling modules: -- prompt_builder – system-prompt composition & context injection -- mcp_builder – MCP server construction & tool-policy helpers -- session_store – on-disk persistence, history, message copying -- agent_loop – the SDK query loop, streaming, mock agent +- agent_manager_ops – edit, close, resume, duplicate, invoke, LLM metadata +- agent_loop – the SDK query loop, streaming, mock agent +- agent_mock – session-completed analytics +- prompt_builder – system-prompt composition & context injection +- mcp_builder – MCP server construction & tool-policy helpers +- session_store – on-disk persistence, history, message copying """ from __future__ import annotations @@ -16,23 +18,24 @@ from datetime import datetime from typing import Optional from uuid import uuid4 -from backend.apps.agents.models import ( - AgentConfig, AgentSession, Message, MessageBranch, ApprovalRequest, ToolGroupMeta, -) +from backend.apps.agents.models import AgentConfig, AgentSession, Message from backend.apps.agents.ws_manager import ws_manager from backend.apps.agents.prompt_builder import resolve_mode from backend.apps.agents.mcp_builder import get_all_tool_names from backend.apps.agents.session_store import ( - save_session, load_session_data, delete_session_file, - load_all_session_data, build_search_text, get_history, + delete_session_file, get_history, reconcile_on_startup, get_browser_agent_children, - copy_session_messages, ) -from backend.apps.agents.agent_loop import ( - run_agent_loop, fire_session_completed, +from backend.apps.agents.agent_loop import run_agent_loop +from backend.apps.agents.agent_manager_ops import ( + edit_message_op, close_session_op, resume_session_op, + duplicate_session_op, invoke_agent_op, +) +from backend.apps.agents.agent_manager_meta import ( + generate_title_op, generate_group_meta_op, + persist_all_sessions_op, restore_all_sessions_op, ) from backend.apps.settings.settings import load_settings -from backend.apps.common.llm_helpers import quick_llm_call, quick_llm_json from backend.apps.analytics.collector import record as _analytics logger = logging.getLogger(__name__) @@ -45,10 +48,6 @@ class AgentManager: self.sessions: dict[str, AgentSession] = {} self.tasks: dict[str, asyncio.Task] = {} - # ------------------------------------------------------------------ - # Session lifecycle - # ------------------------------------------------------------------ - async def launch_agent(self, config: AgentConfig) -> AgentSession: session_id = uuid4().hex mode_tools, _, mode_folder = resolve_mode(config.mode, get_all_tool_names) @@ -60,7 +59,6 @@ class AgentManager: 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( id=session_id, name=config.name, provider=getattr(config, "provider", "anthropic"), @@ -74,10 +72,7 @@ class AgentManager: "model": session.model, "provider": session.provider, "mode": session.mode, "tool_count": len(mode_tools), }, session_id=session_id, dashboard_id=config.dashboard_id) - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "running", - "session": session.model_dump(mode="json"), - }) + await ws_manager.emit_status(session_id, "running", session) return session async def send_message( @@ -112,10 +107,7 @@ class AgentManager: session.allowed_tools = mode_tools session_changed = True if session_changed: - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": session.status, - "session": session.model_dump(mode="json"), - }) + 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 @@ -125,9 +117,7 @@ class AgentManager: forced_tools=forced_tools or None, images=image_meta, hidden=hidden, ) session.messages.append(user_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": user_msg.model_dump(mode="json"), - }) + await ws_manager.emit_message(session_id, user_msg) if context_paths or attached_skills or images or forced_tools: _analytics("context.attached", { "file_count": len([c for c in (context_paths or []) if c.get("type") == "file"]), @@ -146,10 +136,7 @@ class AgentManager: }, session_id=session_id, dashboard_id=session.dashboard_id) session.status = "running" - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "running", - "session": session.model_dump(mode="json"), - }) + 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, @@ -175,10 +162,7 @@ class AgentManager: session.status = "stopped" if not session.closed_at: session.closed_at = datetime.now() - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "stopped", - "session": session.model_dump(mode="json"), - }) + 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) @@ -187,59 +171,7 @@ class AgentManager: ws_manager.resolve_approval(request_id, decision) async def edit_message(self, session_id: str, message_id: str, new_content: str): - 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(): - existing.cancel() - try: - await existing - except asyncio.CancelledError: - pass - - target_msg = next((m for m in session.messages if m.id == message_id), None) - if not target_msg or target_msg.role != "user": - raise ValueError("Can only edit user messages") - - fork_point_id = message_id - fork_parent_branch = target_msg.branch_id - msg_branch = session.branches.get(target_msg.branch_id) - if msg_branch and msg_branch.fork_point_message_id: - branch_user_msgs = [m for m in session.messages if m.branch_id == target_msg.branch_id and m.role == "user"] - if branch_user_msgs and branch_user_msgs[0].id == message_id: - fork_point_id = msg_branch.fork_point_message_id - fork_parent_branch = msg_branch.parent_branch_id or "main" - - new_branch_id = uuid4().hex - new_branch = MessageBranch(id=new_branch_id, parent_branch_id=fork_parent_branch, fork_point_message_id=fork_point_id) - session.branches[new_branch_id] = new_branch - session.active_branch_id = new_branch_id - _analytics("feature.used", { - "feature": "message.branched", - "branch_depth": len([b for b in session.branches.values() if b.parent_branch_id]), - "total_branches_in_session": len(session.branches), - "messages_before_fork": len([m for m in session.messages if m.branch_id == fork_parent_branch]), - }, session_id=session_id, dashboard_id=session.dashboard_id) - - edited_msg = Message( - role="user", content=new_content, branch_id=new_branch_id, - parent_id=target_msg.parent_id, images=target_msg.images, - context_paths=target_msg.context_paths, forced_tools=target_msg.forced_tools, - attached_skills=target_msg.attached_skills, - ) - session.messages.append(edited_msg) - await ws_manager.send_to_session(session_id, "agent:message", {"session_id": session_id, "message": edited_msg.model_dump(mode="json")}) - await ws_manager.send_to_session(session_id, "agent:branch_created", {"session_id": session_id, "branch": new_branch.model_dump(mode="json"), "active_branch_id": new_branch_id}) - session.sdk_session_id = None - session.status = "running" - await ws_manager.send_to_session(session_id, "agent:status", {"session_id": session_id, "status": "running", "session": session.model_dump(mode="json")}) - task = asyncio.create_task(run_agent_loop( - self.sessions, session_id, new_content, - images=target_msg.images, context_paths=target_msg.context_paths, - forced_tools=target_msg.forced_tools, attached_skills=target_msg.attached_skills, - )) - self.tasks[session_id] = task + await edit_message_op(self.sessions, self.tasks, session_id, message_id, new_content) async def switch_branch(self, session_id: str, branch_id: str): session = self.sessions.get(session_id) @@ -248,68 +180,13 @@ class AgentManager: if branch_id not in session.branches: raise ValueError(f"Branch {branch_id} not found") session.active_branch_id = branch_id - await ws_manager.send_to_session(session_id, "agent:branch_switched", {"session_id": session_id, "active_branch_id": branch_id}) - - # ------------------------------------------------------------------ - # LLM-powered metadata - # ------------------------------------------------------------------ + await ws_manager.emit_branch_switched(session_id, branch_id) async def generate_title(self, session_id: str, first_prompt: str) -> str: - session = self.sessions.get(session_id) - if not session: - raise ValueError(f"Session {session_id} not found") - title = first_prompt[:40].strip() - try: - title = await quick_llm_call( - "Generate a concise 3-6 word title for a chat that starts with this message. Return only the title, nothing else.", - first_prompt, max_tokens=30, - ) - title = title.strip('"\'') or first_prompt[:40].strip() - except Exception as e: - logger.warning(f"Title generation failed, using fallback: {e}") - session.name = title - await ws_manager.send_to_session(session_id, "agent:name_updated", {"session_id": session_id, "name": title}) - return title + return await generate_title_op(self.sessions, session_id, first_prompt) async def generate_group_meta(self, session_id: str, group_id: str, tool_calls: list[dict], results_summary: list[str] | None = None, is_refinement: bool = False) -> dict: - session = self.sessions.get(session_id) - if not session: - raise ValueError(f"Session {session_id} not found") - fallback_name = tool_calls[0].get("tool", "Tool calls") if tool_calls else "Tool calls" - fallback_name = fallback_name.split("__")[-1].replace("_", " ").title() if "__" in fallback_name else fallback_name - name, svg = fallback_name, "" - try: - tool_desc = "\n".join(f"- {tc.get('tool', '?')}: {tc.get('input_summary', '')}" for tc in tool_calls) - user_content = f"Tool actions:\n{tool_desc}" - if results_summary: - user_content += "\n\nResults:\n" + "\n".join(f"- {r}" for r in results_summary) - system = ( - "Generate a concise 2-5 word name and a minimal SVG icon for a group of tool actions.\n\n" - "Return ONLY valid JSON: {\"name\": \"...\", \"svg\": \"...\"}\n\n" - "Name rules:\n- 2-5 words, title case, describes the action\n\n" - "SVG rules:\n- 24x24 viewBox\n- Use currentColor for all stroke/fill\n" - "- Simple geometric shapes only\n- No text, no images, no gradients\n" - "- Max 400 characters for the svg string" - ) - parsed = await quick_llm_json(system, user_content) - if parsed.get("name"): - name = parsed["name"].strip().strip("\"'") - if parsed.get("svg"): - svg = parsed["svg"].strip() - except Exception as e: - logger.warning(f"Group meta generation failed, using fallback: {e}") - - meta = ToolGroupMeta(id=group_id, name=name, svg=svg, is_refined=is_refinement) - session.tool_group_meta[group_id] = meta - await ws_manager.send_to_session(session_id, "agent:group_meta_updated", { - "session_id": session_id, "group_id": group_id, - "name": name, "svg": svg, "is_refined": is_refinement, - }) - return {"name": name, "svg": svg, "is_refined": is_refinement} - - # ------------------------------------------------------------------ - # Session management - # ------------------------------------------------------------------ + return await generate_group_meta_op(self.sessions, session_id, group_id, tool_calls, results_summary, is_refinement) async def update_session(self, session_id: str, **fields): session = self.sessions.get(session_id) @@ -318,169 +195,23 @@ class AgentManager: for key, value in fields.items(): if key in {"system_prompt", "name"}: setattr(session, key, value) - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": session.status, - "session": session.model_dump(mode="json"), - }) + await ws_manager.emit_status(session_id, session.status, session) async def close_session(self, session_id: str) -> None: - 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) - task = self.tasks.get(session_id) - if task and not task.done(): - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - session = self.sessions.get(session_id) - if not session: - raise ValueError(f"Session {session_id} not found") - if session.status in ("running", "waiting_approval"): - session.status = "stopped" - session.closed_at = datetime.now() - for req in list(session.pending_approvals): - ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Session closed"}) - session.pending_approvals = [] - if hasattr(session, '_cancel_event'): - session._cancel_event.set() - fire_session_completed(session, self.sessions) - doc_data = session.model_dump(mode="json") - doc_data["search_text"] = build_search_text(session) - save_session(session_id, doc_data) - await ws_manager.send_to_session(session_id, "agent:closed", { - "session_id": session_id, "status": session.status, - "name": session.name, "model": session.model, "mode": session.mode, - "created_at": session.created_at.isoformat() if session.created_at else None, - "closed_at": session.closed_at.isoformat() if session.closed_at else None, - "cost_usd": session.cost_usd, "dashboard_id": session.dashboard_id, - }) - self.sessions.pop(session_id, None) - self.tasks.pop(session_id, None) - logger.info(f"Session {session_id} closed and persisted") + await close_session_op(self.sessions, self.tasks, session_id) async def delete_session(self, session_id: str) -> None: - 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) - task = self.tasks.get(session_id) - if task and not task.done(): - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - self.sessions.pop(session_id, None) - self.tasks.pop(session_id, None) - delete_session_file(session_id) - logger.info(f"Session {session_id} permanently deleted") + from backend.apps.agents.agent_manager_ops import delete_session_op + await delete_session_op(self, session_id) async def resume_session(self, session_id: str) -> AgentSession: - if session_id in self.sessions: - return self.sessions[session_id] - data = load_session_data(session_id) - if data is None: - raise ValueError(f"Session {session_id} not found in history") - session = AgentSession(**data) - hours_since = 0 - if data.get("closed_at"): - try: - closed = datetime.fromisoformat(data["closed_at"][:19]) - hours_since = round((datetime.now() - closed).total_seconds() / 3600, 1) - except Exception: - pass - _analytics("session.resumed", { - "hours_since_closed": hours_since, - "original_message_count": len(data.get("messages", [])), - "original_cost_usd": data.get("cost_usd", 0), "model": session.model, - }, session_id=session_id, dashboard_id=session.dashboard_id) - session.closed_at = None - self.sessions[session_id] = session - delete_session_file(session_id) - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": session.status, - "session": session.model_dump(mode="json"), - }) - logger.info(f"Session {session_id} resumed from history") - return session + return await resume_session_op(self.sessions, session_id) async def duplicate_session(self, session_id: str, dashboard_id: str | None = None, up_to_message_id: str | None = None) -> AgentSession: - source = self.sessions.get(session_id) - if not source: - data = load_session_data(session_id) - if data is None: - raise ValueError(f"Session {session_id} not found") - source = AgentSession(**data) - new_messages, new_branches, _ = copy_session_messages(source, up_to_message_id) - new_session = AgentSession( - id=uuid4().hex, name=f"{source.name} (copy)", status="stopped", - model=source.model, mode=source.mode, system_prompt=source.system_prompt, - allowed_tools=list(source.allowed_tools), max_turns=source.max_turns, - cwd=source.cwd, created_at=datetime.now(), messages=new_messages, - branches=new_branches, active_branch_id=source.active_branch_id, - tool_group_meta=dict(source.tool_group_meta), - dashboard_id=dashboard_id or source.dashboard_id, - ) - self.sessions[new_session.id] = new_session - await ws_manager.send_to_session(new_session.id, "agent:status", { - "session_id": new_session.id, "status": new_session.status, - "session": new_session.model_dump(mode="json"), - }) - return new_session + return await duplicate_session_op(self.sessions, session_id, dashboard_id, up_to_message_id) async def invoke_agent(self, source_session_id: str, message: str, parent_session_id: str | None = None, dashboard_id: str | None = None) -> dict: - source = self.sessions.get(source_session_id) - if not source: - data = load_session_data(source_session_id) - if data is None: - raise ValueError(f"Session {source_session_id} not found") - source = AgentSession(**data) - source_name = source.name - new_messages, new_branches, _ = copy_session_messages(source) - fork = AgentSession( - id=uuid4().hex, name=f"{source_name} (invoked)", status="running", - model=source.model, mode="invoked-agent", sdk_session_id=source.sdk_session_id, - system_prompt=source.system_prompt, allowed_tools=list(source.allowed_tools), - max_turns=source.max_turns or 25, cwd=source.cwd, created_at=datetime.now(), - messages=new_messages, branches=new_branches, - active_branch_id=source.active_branch_id, - tool_group_meta=dict(source.tool_group_meta), - dashboard_id=dashboard_id or source.dashboard_id, - parent_session_id=parent_session_id, - ) - self.sessions[fork.id] = fork - await ws_manager.broadcast_global("agent:status", { - "session_id": fork.id, "status": fork.status, - "session": fork.model_dump(mode="json"), - }) - user_msg = Message(role="user", content=message, branch_id=fork.active_branch_id) - fork.messages.append(user_msg) - await ws_manager.send_to_session(fork.id, "agent:message", { - "session_id": fork.id, "message": user_msg.model_dump(mode="json"), - }) - await run_agent_loop(self.sessions, fork.id, message, fork_session=True) - last_assistant = None - for msg in reversed(fork.messages): - if msg.role == "assistant": - content = msg.content - if isinstance(content, str): - last_assistant = content - elif isinstance(content, list): - texts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"] - last_assistant = "\n".join(texts) - else: - last_assistant = str(content) - break - return { - "forked_session_id": fork.id, "source_name": source_name, - "response": last_assistant or "No response from invoked agent.", - "cost_usd": fork.cost_usd, - } - - # ------------------------------------------------------------------ - # Queries - # ------------------------------------------------------------------ + return await invoke_agent_op(self.sessions, source_session_id, message, parent_session_id, dashboard_id) def get_all_sessions(self, dashboard_id: str | None = None) -> list[AgentSession]: if dashboard_id: @@ -490,10 +221,6 @@ class AgentManager: def get_session(self, session_id: str) -> Optional[AgentSession]: return self.sessions.get(session_id) - # ------------------------------------------------------------------ - # Delegated helpers (kept as methods for API compatibility) - # ------------------------------------------------------------------ - def get_history(self, q: str = "", limit: int = 20, offset: int = 0, dashboard_id: str | None = None) -> dict: return get_history(q=q, limit=limit, offset=offset, dashboard_id=dashboard_id) @@ -501,35 +228,10 @@ class AgentManager: return await reconcile_on_startup() async def persist_all_sessions(self) -> None: - for session_id, session in list(self.sessions.items()): - if session.status in ("running", "waiting_approval"): - session.status = "stopped" - for req in list(session.pending_approvals): - ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Server shutting down"}) - session.pending_approvals = [] - fire_session_completed(session, self.sessions) - doc_data = session.model_dump(mode="json") - doc_data["search_text"] = build_search_text(session) - save_session(session_id, doc_data) - logger.info(f"Persisted session {session_id} on shutdown") - self.sessions.clear() - self.tasks.clear() + await persist_all_sessions_op(self.sessions, self.tasks) async def restore_all_sessions(self) -> None: - for sid, data in load_all_session_data(): - try: - session = AgentSession(**data) - except Exception as e: - logger.warning(f"Skipping corrupt session file {sid}: {e}") - continue - if session.closed_at is not None: - continue - if session.status in ("running", "waiting_approval"): - session.status = "stopped" - session.pending_approvals = [] - self.sessions[session.id] = session - delete_session_file(sid) - logger.info(f"Restored session {session.id}") + await restore_all_sessions_op(self.sessions) def get_browser_agent_children(self, parent_session_id: str) -> list[dict]: return get_browser_agent_children(self.sessions, parent_session_id) diff --git a/backend/apps/agents/agent_manager_meta.py b/backend/apps/agents/agent_manager_meta.py new file mode 100644 index 00000000..673d70be --- /dev/null +++ b/backend/apps/agents/agent_manager_meta.py @@ -0,0 +1,126 @@ +"""LLM-powered metadata, persistence helpers, and delete operation. + +Extracted from agent_manager_ops to keep every file under 250 lines. +""" + +from __future__ import annotations + +import asyncio +import logging + +from backend.apps.agents.models import AgentSession, ToolGroupMeta +from backend.apps.agents.ws_manager import ws_manager +from backend.apps.agents.session_store import ( + save_session, delete_session_file, build_search_text, +) +from backend.apps.common.llm_helpers import quick_llm_call, quick_llm_json + +logger = logging.getLogger(__name__) + + +async def generate_title_op(sessions: dict, session_id: str, first_prompt: str) -> str: + session = sessions.get(session_id) + if not session: + raise ValueError(f"Session {session_id} not found") + title = first_prompt[:40].strip() + try: + title = await quick_llm_call( + "Generate a concise 3-6 word title for a chat that starts with this message. Return only the title, nothing else.", + first_prompt, max_tokens=30, + ) + title = title.strip('"\'') or first_prompt[:40].strip() + except Exception as e: + logger.warning(f"Title generation failed, using fallback: {e}") + session.name = title + await ws_manager.emit_name_updated(session_id, title) + return title + + +async def generate_group_meta_op( + sessions: dict, session_id: str, group_id: str, + tool_calls: list[dict], results_summary: list[str] | None = None, + is_refinement: bool = False, +) -> dict: + session = sessions.get(session_id) + if not session: + raise ValueError(f"Session {session_id} not found") + fallback_name = tool_calls[0].get("tool", "Tool calls") if tool_calls else "Tool calls" + fallback_name = fallback_name.split("__")[-1].replace("_", " ").title() if "__" in fallback_name else fallback_name + name, svg = fallback_name, "" + try: + tool_desc = "\n".join(f"- {tc.get('tool', '?')}: {tc.get('input_summary', '')}" for tc in tool_calls) + user_content = f"Tool actions:\n{tool_desc}" + if results_summary: + user_content += "\n\nResults:\n" + "\n".join(f"- {r}" for r in results_summary) + system = ( + "Generate a concise 2-5 word name and a minimal SVG icon for a group of tool actions.\n\n" + "Return ONLY valid JSON: {\"name\": \"...\", \"svg\": \"...\"}\n\n" + "Name rules:\n- 2-5 words, title case, describes the action\n\n" + "SVG rules:\n- 24x24 viewBox\n- Use currentColor for all stroke/fill\n" + "- Simple geometric shapes only\n- No text, no images, no gradients\n" + "- Max 400 characters for the svg string" + ) + parsed = await quick_llm_json(system, user_content) + if parsed.get("name"): + name = parsed["name"].strip().strip("\"'") + if parsed.get("svg"): + svg = parsed["svg"].strip() + except Exception as e: + logger.warning(f"Group meta generation failed, using fallback: {e}") + + meta = ToolGroupMeta(id=group_id, name=name, svg=svg, is_refined=is_refinement) + session.tool_group_meta[group_id] = meta + await ws_manager.emit_group_meta_updated(session_id, group_id, name, svg, is_refinement) + return {"name": name, "svg": svg, "is_refined": is_refinement} + + +async def persist_all_sessions_op(sessions: dict, tasks: dict) -> None: + from backend.apps.agents.agent_mock import fire_session_completed + for session_id, session in list(sessions.items()): + if session.status in ("running", "waiting_approval"): + session.status = "stopped" + for req in list(session.pending_approvals): + ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Server shutting down"}) + session.pending_approvals = [] + fire_session_completed(session, sessions) + doc_data = session.model_dump(mode="json") + doc_data["search_text"] = build_search_text(session) + save_session(session_id, doc_data) + logger.info(f"Persisted session {session_id} on shutdown") + sessions.clear() + tasks.clear() + + +async def restore_all_sessions_op(sessions: dict) -> None: + from backend.apps.agents.session_store import load_all_session_data + for sid, data in load_all_session_data(): + try: + session = AgentSession(**data) + except Exception as e: + logger.warning(f"Skipping corrupt session file {sid}: {e}") + continue + if session.closed_at is not None: + continue + if session.status in ("running", "waiting_approval"): + session.status = "stopped" + session.pending_approvals = [] + sessions[session.id] = session + delete_session_file(sid) + logger.info(f"Restored session {session.id}") + + +async def delete_session_op(manager, session_id: str) -> None: + children = [s for s in manager.sessions.values() if s.parent_session_id == session_id and s.mode == "browser-agent"] + for child in children: + await manager.stop_agent(child.id) + task = manager.tasks.get(session_id) + if task and not task.done(): + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + manager.sessions.pop(session_id, None) + manager.tasks.pop(session_id, None) + delete_session_file(session_id) + logger.info(f"Session {session_id} permanently deleted") diff --git a/backend/apps/agents/agent_manager_ops.py b/backend/apps/agents/agent_manager_ops.py new file mode 100644 index 00000000..e0ccb552 --- /dev/null +++ b/backend/apps/agents/agent_manager_ops.py @@ -0,0 +1,232 @@ +"""Complex agent-manager operations extracted for the 250-line limit. + +Each function is a standalone async operation that receives the sessions +dict (and other dependencies) explicitly. +""" + +from __future__ import annotations + +import asyncio +import logging +from datetime import datetime +from uuid import uuid4 + +from backend.apps.agents.models import ( + AgentSession, Message, MessageBranch, +) +from backend.apps.agents.ws_manager import ws_manager +from backend.apps.agents.session_store import ( + save_session, load_session_data, delete_session_file, + build_search_text, copy_session_messages, +) +from backend.apps.analytics.collector import record as _analytics + +logger = logging.getLogger(__name__) + + +async def edit_message_op( + sessions: dict, tasks: dict, + session_id: str, message_id: str, new_content: str, +): + session = sessions.get(session_id) + if not session: + raise ValueError(f"Session {session_id} not found") + existing = tasks.get(session_id) + if existing and not existing.done(): + existing.cancel() + try: + await existing + except asyncio.CancelledError: + pass + + target_msg = next((m for m in session.messages if m.id == message_id), None) + if not target_msg or target_msg.role != "user": + raise ValueError("Can only edit user messages") + + fork_point_id = message_id + fork_parent_branch = target_msg.branch_id + msg_branch = session.branches.get(target_msg.branch_id) + if msg_branch and msg_branch.fork_point_message_id: + branch_user_msgs = [m for m in session.messages if m.branch_id == target_msg.branch_id and m.role == "user"] + if branch_user_msgs and branch_user_msgs[0].id == message_id: + fork_point_id = msg_branch.fork_point_message_id + fork_parent_branch = msg_branch.parent_branch_id or "main" + + new_branch_id = uuid4().hex + new_branch = MessageBranch(id=new_branch_id, parent_branch_id=fork_parent_branch, fork_point_message_id=fork_point_id) + session.branches[new_branch_id] = new_branch + session.active_branch_id = new_branch_id + _analytics("feature.used", { + "feature": "message.branched", + "branch_depth": len([b for b in session.branches.values() if b.parent_branch_id]), + "total_branches_in_session": len(session.branches), + "messages_before_fork": len([m for m in session.messages if m.branch_id == fork_parent_branch]), + }, session_id=session_id, dashboard_id=session.dashboard_id) + + edited_msg = Message( + role="user", content=new_content, branch_id=new_branch_id, + parent_id=target_msg.parent_id, images=target_msg.images, + context_paths=target_msg.context_paths, forced_tools=target_msg.forced_tools, + attached_skills=target_msg.attached_skills, + ) + session.messages.append(edited_msg) + await ws_manager.emit_message(session_id, edited_msg) + await ws_manager.emit_branch_created(session_id, new_branch, new_branch_id) + session.sdk_session_id = None + session.status = "running" + await ws_manager.emit_status(session_id, "running", session) + from backend.apps.agents.agent_loop import run_agent_loop + task = asyncio.create_task(run_agent_loop( + sessions, session_id, new_content, + images=target_msg.images, context_paths=target_msg.context_paths, + forced_tools=target_msg.forced_tools, attached_skills=target_msg.attached_skills, + )) + tasks[session_id] = task + + +async def close_session_op( + sessions: dict, tasks: dict, + session_id: str, +): + from backend.apps.agents.agent_manager import agent_manager + children = [s for s in sessions.values() if s.parent_session_id == session_id and s.mode == "browser-agent"] + for child in children: + await agent_manager.stop_agent(child.id) + task = tasks.get(session_id) + if task and not task.done(): + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + session = sessions.get(session_id) + if not session: + raise ValueError(f"Session {session_id} not found") + if session.status in ("running", "waiting_approval"): + session.status = "stopped" + session.closed_at = datetime.now() + for req in list(session.pending_approvals): + ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Session closed"}) + session.pending_approvals = [] + if hasattr(session, '_cancel_event'): + session._cancel_event.set() + from backend.apps.agents.agent_mock import fire_session_completed + fire_session_completed(session, sessions) + doc_data = session.model_dump(mode="json") + doc_data["search_text"] = build_search_text(session) + save_session(session_id, doc_data) + await ws_manager.emit_closed(session_id, session) + sessions.pop(session_id, None) + tasks.pop(session_id, None) + logger.info(f"Session {session_id} closed and persisted") + + +async def resume_session_op(sessions: dict, session_id: str) -> AgentSession: + if session_id in sessions: + return sessions[session_id] + data = load_session_data(session_id) + if data is None: + raise ValueError(f"Session {session_id} not found in history") + session = AgentSession(**data) + hours_since = 0 + if data.get("closed_at"): + try: + closed = datetime.fromisoformat(data["closed_at"][:19]) + hours_since = round((datetime.now() - closed).total_seconds() / 3600, 1) + except Exception: + pass + _analytics("session.resumed", { + "hours_since_closed": hours_since, + "original_message_count": len(data.get("messages", [])), + "original_cost_usd": data.get("cost_usd", 0), "model": session.model, + }, session_id=session_id, dashboard_id=session.dashboard_id) + session.closed_at = None + sessions[session_id] = session + delete_session_file(session_id) + await ws_manager.emit_status(session_id, session.status, session) + logger.info(f"Session {session_id} resumed from history") + return session + + +async def duplicate_session_op( + sessions: dict, session_id: str, + dashboard_id: str | None = None, up_to_message_id: str | None = None, +) -> AgentSession: + source = sessions.get(session_id) + if not source: + data = load_session_data(session_id) + if data is None: + raise ValueError(f"Session {session_id} not found") + source = AgentSession(**data) + new_messages, new_branches, _ = copy_session_messages(source, up_to_message_id) + new_session = AgentSession( + id=uuid4().hex, name=f"{source.name} (copy)", status="stopped", + model=source.model, mode=source.mode, system_prompt=source.system_prompt, + allowed_tools=list(source.allowed_tools), max_turns=source.max_turns, + cwd=source.cwd, created_at=datetime.now(), messages=new_messages, + branches=new_branches, active_branch_id=source.active_branch_id, + tool_group_meta=dict(source.tool_group_meta), + dashboard_id=dashboard_id or source.dashboard_id, + ) + sessions[new_session.id] = new_session + await ws_manager.emit_status(new_session.id, new_session.status, new_session) + return new_session + + +async def invoke_agent_op( + sessions: dict, source_session_id: str, message: str, + parent_session_id: str | None = None, dashboard_id: str | None = None, +) -> dict: + source = sessions.get(source_session_id) + if not source: + data = load_session_data(source_session_id) + if data is None: + raise ValueError(f"Session {source_session_id} not found") + source = AgentSession(**data) + source_name = source.name + new_messages, new_branches, _ = copy_session_messages(source) + fork = AgentSession( + id=uuid4().hex, name=f"{source_name} (invoked)", status="running", + model=source.model, mode="invoked-agent", sdk_session_id=source.sdk_session_id, + system_prompt=source.system_prompt, allowed_tools=list(source.allowed_tools), + max_turns=source.max_turns or 25, cwd=source.cwd, created_at=datetime.now(), + messages=new_messages, branches=new_branches, + active_branch_id=source.active_branch_id, + tool_group_meta=dict(source.tool_group_meta), + dashboard_id=dashboard_id or source.dashboard_id, + parent_session_id=parent_session_id, + ) + sessions[fork.id] = fork + await ws_manager.broadcast_global("agent:status", { + "session_id": fork.id, "status": fork.status, + "session": fork.model_dump(mode="json"), + }) + user_msg = Message(role="user", content=message, branch_id=fork.active_branch_id) + fork.messages.append(user_msg) + await ws_manager.emit_message(fork.id, user_msg) + from backend.apps.agents.agent_loop import run_agent_loop + await run_agent_loop(sessions, fork.id, message, fork_session=True) + last_assistant = None + for msg in reversed(fork.messages): + if msg.role == "assistant": + content = msg.content + if isinstance(content, str): + last_assistant = content + elif isinstance(content, list): + texts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"] + last_assistant = "\n".join(texts) + else: + last_assistant = str(content) + break + return { + "forked_session_id": fork.id, "source_name": source_name, + "response": last_assistant or "No response from invoked agent.", + "cost_usd": fork.cost_usd, + } + + +from backend.apps.agents.agent_manager_meta import ( # noqa: F401 — re-exports + generate_title_op, generate_group_meta_op, + persist_all_sessions_op, restore_all_sessions_op, + delete_session_op, +) diff --git a/backend/apps/agents/agent_mock.py b/backend/apps/agents/agent_mock.py new file mode 100644 index 00000000..774f9265 --- /dev/null +++ b/backend/apps/agents/agent_mock.py @@ -0,0 +1,132 @@ +"""Mock agent and session-completed analytics. + +Extracted from agent_loop.py to keep every file under 250 lines. +""" + +from __future__ import annotations + +import asyncio +import json +import logging +from datetime import datetime +from uuid import uuid4 + +from backend.apps.agents.models import AgentSession, ApprovalRequest, Message +from backend.apps.agents.ws_manager import ws_manager +from backend.apps.analytics.collector import record as _analytics + +logger = logging.getLogger(__name__) + + +async def stream_text(session_id: str, msg_id: str, text: str, delay: float = 0.03): + await ws_manager.emit_stream_start(session_id, msg_id, "assistant") + words = text.split(" ") + for i, word in enumerate(words): + chunk = word if i == 0 else " " + word + await ws_manager.emit_stream_delta(session_id, msg_id, chunk) + await asyncio.sleep(delay) + await ws_manager.emit_stream_end(session_id, msg_id) + + +async def stream_tool_input(session_id: str, msg_id: str, tool_name: str, input_json: str, delay: float = 0.02): + await ws_manager.emit_stream_start(session_id, msg_id, "tool_call", tool_name=tool_name) + chunk_size = 12 + for i in range(0, len(input_json), chunk_size): + await ws_manager.emit_stream_delta(session_id, msg_id, input_json[i:i + chunk_size]) + await asyncio.sleep(delay) + await ws_manager.emit_stream_end(session_id, msg_id) + + +def fire_session_completed(session: AgentSession, sessions_dict: dict[str, AgentSession]): + duration = 0.0 + if session.created_at: + end = session.closed_at or datetime.now() + duration = (end - session.created_at).total_seconds() + tool_names = [ + m.content.get("tool", "") for m in session.messages + if m.role == "tool_call" and isinstance(m.content, dict) + ] + user_messages = [ + (m.content if isinstance(m.content, str) else str(m.content))[:200] + for m in session.messages if m.role == "user" + ] + _analytics("session.completed", { + "model": session.model, + "provider": getattr(session, "provider", "anthropic"), + "mode": session.mode, + "cost_usd": session.cost_usd, + "message_count": len([m for m in session.messages if m.role in ("user", "assistant")]), + "duration_seconds": round(duration, 1), + "status": session.status, + "tool_count": len(tool_names), + "tools_list": list(set(tool_names)), + "session_title": session.name, + "first_user_message": user_messages[0] if user_messages else "", + "input_tokens": session.tokens.get("input", 0), + "output_tokens": session.tokens.get("output", 0), + "is_sub_agent": session.parent_session_id is not None, + "parent_session_id": session.parent_session_id, + "sub_agent_count": len([s for s in sessions_dict.values() if s.parent_session_id == session.id]), + "branch_count": len(session.branches), + }, session_id=session.id, dashboard_id=session.dashboard_id) + + +async def run_mock_agent(session_id: str, prompt: str, sessions: dict[str, AgentSession]): + session = sessions.get(session_id) + if not session: + return + + await asyncio.sleep(1) + + request_id = uuid4().hex + approval_req = ApprovalRequest( + id=request_id, session_id=session_id, tool_name="Bash", + tool_input={"command": f"echo 'Processing: {prompt}'", "description": "Echo the user prompt"}, + ) + session.pending_approvals.append(approval_req) + session.status = "waiting_approval" + await ws_manager.emit_status(session_id, "waiting_approval") + + decision = await ws_manager.send_approval_request( + session_id, request_id, "Bash", + {"command": f"echo 'Processing: {prompt}'", "description": "Echo the user prompt"}, + ) + + session.pending_approvals = [a for a in session.pending_approvals if a.id != request_id] + session.status = "running" + await ws_manager.emit_status(session_id, "running") + + tool_input_content = {"tool": "Bash", "input": {"command": f"echo 'Processing: {prompt}'"}, "approved": decision.get("behavior") == "allow"} + tool_msg_id = uuid4().hex + await stream_tool_input(session_id, tool_msg_id, "Bash", json.dumps(tool_input_content["input"], indent=2)) + tool_msg = Message(id=tool_msg_id, role="tool_call", content=tool_input_content, branch_id=session.active_branch_id) + session.messages.append(tool_msg) + await ws_manager.emit_message(session_id, tool_msg) + + await asyncio.sleep(1) + + if decision.get("behavior") == "allow": + tool_result = Message(role="tool_result", content=f"Processing: {prompt}", branch_id=session.active_branch_id) + session.messages.append(tool_result) + await ws_manager.emit_message(session_id, tool_result) + + await asyncio.sleep(1) + + asst_text = ( + f"I've processed your request: \"{prompt}\"\n\n" + "This is a mock response because `claude-agent-sdk` is not installed. " + "Install it with `pip install claude-agent-sdk` to use real Claude Code instances.\n\n" + f"The agent was configured with:\n- Model: {session.model}\n- Mode: {session.mode}" + ) + asst_msg_id = uuid4().hex + await stream_text(session_id, asst_msg_id, asst_text) + + asst_msg = Message(id=asst_msg_id, role="assistant", content=asst_text, branch_id=session.active_branch_id) + session.messages.append(asst_msg) + await ws_manager.emit_message(session_id, asst_msg) + + session.status = "completed" + session.closed_at = datetime.now() + session.cost_usd = 0.001 + await ws_manager.emit_status(session_id, "completed", session) + await ws_manager.emit_cost_update(session_id, session.cost_usd) diff --git a/backend/apps/agents/agent_options.py b/backend/apps/agents/agent_options.py new file mode 100644 index 00000000..a043b705 --- /dev/null +++ b/backend/apps/agents/agent_options.py @@ -0,0 +1,217 @@ +"""Build ClaudeAgentOptions kwargs and MCP server configuration. + +Extracted from agent_loop.py to keep every file under 250 lines. +""" + +from __future__ import annotations + +import json +import logging +import os +import sys + +from backend.apps.agents.models import AgentSession +from backend.apps.agents.prompt_builder import ( + resolve_mode, compose_system_prompt, build_connected_tools_context, + build_outputs_context, build_browser_context, get_pre_selected_browser_ids, +) +from backend.apps.agents.mcp_builder import ( + FULL_TOOLS, build_mcp_servers, get_all_tool_names, + _get_denied_tool_names, _get_all_known_tool_names, _is_fully_denied, +) +from backend.apps.settings.settings import load_settings +from backend.apps.tools_lib.tools_lib import ( + _load_all as load_all_tools, + load_builtin_permissions, +) +from backend.apps.common.mcp_utils import sanitize_server_name as _sanitize_server_name + +logger = logging.getLogger(__name__) + + +async def build_agent_options( + session: AgentSession, + builtin_perms: dict, + can_use_tool, + pre_tool_hook, + post_tool_hook, + fork_session: bool = False, + selected_browser_ids: list[str] | None = None, +) -> dict: + """Build the kwargs dict for ClaudeAgentOptions. + + Requires claude_agent_sdk types to be imported by the caller; they are + passed in via the hook callables. + """ + from claude_agent_sdk import ClaudeAgentOptions + from claude_agent_sdk.types import HookMatcher + + _, mode_sys_prompt, _ = resolve_mode(session.mode, get_all_tool_names) + connected_tools_ctx = build_connected_tools_context( + session.allowed_tools, load_all_tools, get_all_tool_names, _is_fully_denied, _get_denied_tool_names, + ) + outputs_ctx = build_outputs_context() + browser_ctx = build_browser_context(session.dashboard_id, selected_browser_ids=selected_browser_ids) + global_settings = load_settings() + composed_prompt = compose_system_prompt( + global_settings.default_system_prompt, mode_sys_prompt, session.system_prompt, + connected_tools_ctx, outputs_ctx, browser_ctx, + ) + + if session.mode == "view-builder": + from backend.apps.outputs.view_builder_templates import VIEW_BUILDER_SKILL + skill_block = f"\n{VIEW_BUILDER_SKILL}\n" + composed_prompt = f"{composed_prompt}\n\n{skill_block}" if composed_prompt else skill_block + + mcp_servers = await build_mcp_servers(session.allowed_tools) + + _browser_delegation_tools = ["CreateBrowserAgent", "BrowserAgent", "BrowserAgents"] + _browser_all_denied = all(builtin_perms.get(t, "always_allow") == "deny" for t in _browser_delegation_tools) + + if not _browser_all_denied: + browser_agent_server_path = os.path.join(os.path.dirname(__file__), "browser_agent_mcp_server.py") + backend_port = os.environ.get("OPENSWARM_PORT", "8324") + pre_selected_bids = get_pre_selected_browser_ids(session.dashboard_id) + mcp_servers["openswarm-browser-agent"] = { + "command": sys.executable, + "args": [browser_agent_server_path], + "env": { + "OPENSWARM_PORT": backend_port, + "OPENSWARM_AGENT_MODEL": session.model, + "OPENSWARM_DASHBOARD_ID": session.dashboard_id or "", + "OPENSWARM_PRE_SELECTED_BROWSER_IDS": ",".join(pre_selected_bids), + "OPENSWARM_PARENT_SESSION_ID": session.id, + }, + "type": "stdio", + } + + _invoke_agent_tools = ["InvokeAgent"] + _invoke_all_denied = all(builtin_perms.get(t, "always_allow") == "deny" for t in _invoke_agent_tools) + + if not _invoke_all_denied: + invoke_agent_server_path = os.path.join(os.path.dirname(__file__), "invoke_agent_mcp_server.py") + backend_port = os.environ.get("OPENSWARM_PORT", "8324") + mcp_servers["openswarm-invoke-agent"] = { + "command": sys.executable, + "args": [invoke_agent_server_path], + "env": { + "OPENSWARM_PORT": backend_port, + "OPENSWARM_PARENT_SESSION_ID": session.id, + "OPENSWARM_DASHBOARD_ID": session.dashboard_id or "", + }, + "type": "stdio", + } + + effective_allowed, effective_disallowed = _compute_tool_permissions( + session, builtin_perms, mcp_servers, _browser_delegation_tools, _invoke_agent_tools, + ) + + options_kwargs: dict = { + "model": session.model, + "max_buffer_size": 5 * 1024 * 1024, + "permission_mode": "default", + "can_use_tool": can_use_tool, + "hooks": { + "PreToolUse": [HookMatcher(matcher=None, hooks=[pre_tool_hook])], + "PostToolUse": [HookMatcher(matcher=None, hooks=[post_tool_hook])], + }, + "allowed_tools": effective_allowed, + "disallowed_tools": effective_disallowed, + "include_partial_messages": True, + } + + from backend.apps.nine_router import is_running as _9r_running + if global_settings.anthropic_api_key: + options_kwargs["env"] = {"ANTHROPIC_API_KEY": global_settings.anthropic_api_key} + logger.info("[MCP-DEBUG] Using direct API key") + elif _9r_running(): + options_kwargs["env"] = { + "ANTHROPIC_API_KEY": "9router", + "ANTHROPIC_BASE_URL": "http://localhost:20128", + } + options_kwargs["extra_args"] = {"bare": None} + logger.info("[MCP-DEBUG] Using 9Router (bare mode)") + else: + raise ValueError("No AI provider configured. Set an API key or connect a subscription.") + + if mcp_servers: + options_kwargs["mcp_servers"] = mcp_servers + mcp_json_len = len(json.dumps({"mcpServers": mcp_servers})) + logger.info(f"[MCP-DEBUG] mcp_servers passed to SDK: {list(mcp_servers.keys())}, JSON length={mcp_json_len}") + if composed_prompt: + options_kwargs["system_prompt"] = composed_prompt + if session.max_turns: + options_kwargs["max_turns"] = session.max_turns + if session.cwd: + options_kwargs["cwd"] = session.cwd + if session.sdk_session_id: + options_kwargs["resume"] = session.sdk_session_id + if fork_session: + options_kwargs["fork_session"] = True + + logger.info(f"[MCP-DEBUG] Creating ClaudeAgentOptions with model={session.model}") + return options_kwargs + + +def _compute_tool_permissions( + session: AgentSession, + builtin_perms: dict, + mcp_servers: dict, + browser_delegation_tools: list[str], + invoke_agent_tools: list[str], +) -> tuple[list[str], list[str]]: + effective_allowed = [ + t for t in session.allowed_tools + if t in FULL_TOOLS and builtin_perms.get(t, "always_allow") == "always_allow" + ] + effective_disallowed = [ + t for t in FULL_TOOLS + if builtin_perms.get(t, "always_allow") == "deny" + ] + + if not mcp_servers: + return effective_allowed, effective_disallowed + + all_tools_list = load_all_tools() + for name in mcp_servers: + if name == "openswarm-browser-agent": + for bt in browser_delegation_tools: + policy = builtin_perms.get(bt, "always_allow") + if policy == "always_allow": + effective_allowed.append(f"mcp__openswarm-browser-agent__{bt}") + elif policy == "deny": + effective_disallowed.append(f"mcp__openswarm-browser-agent__{bt}") + continue + if name == "openswarm-invoke-agent": + for it in invoke_agent_tools: + policy = builtin_perms.get(it, "always_allow") + if policy == "always_allow": + effective_allowed.append(f"mcp__openswarm-invoke-agent__{it}") + elif policy == "deny": + effective_disallowed.append(f"mcp__openswarm-invoke-agent__{it}") + continue + tool_def = next( + (t for t in all_tools_list if t.mcp_config and t.enabled and _sanitize_server_name(t.name) == name), + None, + ) + if tool_def: + denied = _get_denied_tool_names(tool_def) + known = _get_all_known_tool_names(tool_def) + for tn in known - denied: + policy = tool_def.tool_permissions.get(tn, "ask") + if policy == "always_allow": + effective_allowed.append(f"mcp__{name}__{tn}") + for tn in denied: + effective_disallowed.append(f"mcp__{name}__{tn}") + else: + effective_allowed.append(f"mcp__{name}__*") + + google_allowed = [t for t in effective_allowed if "google-workspace" in t] + reddit_allowed = [t for t in effective_allowed if "reddit" in t] + builtin_allowed = [t for t in effective_allowed if not t.startswith("mcp__")] + logger.info(f"[MCP-DEBUG] effective_allowed: {len(effective_allowed)} total " + f"(builtins={len(builtin_allowed)}, google={len(google_allowed)}, reddit={len(reddit_allowed)})") + if effective_disallowed: + logger.info(f"[MCP-DEBUG] effective_disallowed: {effective_disallowed}") + + return effective_allowed, effective_disallowed diff --git a/backend/apps/agents/agents.py b/backend/apps/agents/agents.py index fbd14737..996f9604 100644 --- a/backend/apps/agents/agents.py +++ b/backend/apps/agents/agents.py @@ -3,9 +3,9 @@ from backend.apps.agents.agent_manager import agent_manager from backend.apps.agents.ws_manager import ws_manager from backend.apps.agents.models import AgentConfig, ApprovalResponse from contextlib import asynccontextmanager -from fastapi import WebSocket, WebSocketDisconnect, HTTPException +from fastapi import HTTPException, Request from fastapi.responses import JSONResponse -import json +from uuid import uuid4 import logging logger = logging.getLogger(__name__) @@ -179,127 +179,57 @@ async def resume_session(session_id: str): return {"session": session.model_dump(mode="json")} -# --------------------------------------------------------------------------- -# 9Router / Subscription endpoints -# --------------------------------------------------------------------------- - -@agents.router.get("/subscriptions/status") -async def subscriptions_status(): - """Check if 9Router is running and list connected providers.""" - from backend.apps.nine_router import is_running, get_providers, get_models - if not is_running(): - return {"running": False, "providers": [], "models": []} - providers = await get_providers() - models = await get_models() - return {"running": True, "providers": providers, "models": models} +@agents.router.post("/browser/command") +async def browser_command(request: Request): + """Proxy browser commands to the frontend via WebSocket and wait for results.""" + body = await request.json() + action = body.get("action", "") + browser_id = body.get("browser_id", "") + tab_id = body.get("tab_id", "") + params = body.get("params", {}) + if not action or not browser_id: + return JSONResponse({"error": "action and browser_id are required"}, status_code=400) + request_id = uuid4().hex + result = await ws_manager.send_browser_command(request_id, action, browser_id, params, tab_id=tab_id) + return JSONResponse(result) -@agents.router.post("/subscriptions/connect") -async def subscriptions_connect(body: dict): - """Start OAuth flow for a subscription provider.""" - from backend.apps.nine_router import is_running, ensure_running, start_oauth - provider = body.get("provider", "") - if not provider: - raise HTTPException(status_code=400, detail="provider required") +@agents.router.post("/browser-agent/run") +async def browser_agent_run(request: Request): + """Run one or more browser sub-agents in parallel.""" + from backend.apps.agents.browser_agent import run_browser_agents + body = await request.json() + tasks = body.get("tasks", []) + if not tasks: + return JSONResponse({"error": "tasks array is required"}, status_code=400) + results = await run_browser_agents( + tasks=tasks, model=body.get("model", "sonnet"), + dashboard_id=body.get("dashboard_id", "") or None, + pre_selected_browser_ids=body.get("pre_selected_browser_ids", []), + parent_session_id=body.get("parent_session_id", "") or None, + ) + return JSONResponse({"results": results}) - if not is_running(): - await ensure_running() - if not is_running(): - raise HTTPException(status_code=503, detail="9Router not available. Please install Node.js.") +@agents.router.post("/invoke-agent/run") +async def invoke_agent_run(request: Request): + """Fork an existing agent session and send it a new message.""" + body = await request.json() + session_id = body.get("session_id", "") + message = body.get("message", "") + if not session_id: + return JSONResponse({"error": "session_id is required"}, status_code=400) + if not message: + return JSONResponse({"error": "message is required"}, status_code=400) try: - result = await start_oauth(provider) - - # For auth_code flows, store pending state so the callback can exchange - if result.get("flow") == "authorization_code" and result.get("state"): - from backend.main import _pending_oauth - _pending_oauth[result["state"]] = { - "provider": provider, - "code_verifier": result.get("code_verifier", ""), - "redirect_uri": result.get("redirect_uri", ""), - } - - return result - except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) - - -@agents.router.post("/subscriptions/poll") -async def subscriptions_poll(body: dict): - """Poll for OAuth completion.""" - from backend.apps.nine_router import poll_oauth - provider = body.get("provider", "") - device_code = body.get("device_code", "") - if not provider or not device_code: - raise HTTPException(status_code=400, detail="provider and device_code required") - - try: - result = await poll_oauth( - provider, device_code, - code_verifier=body.get("code_verifier"), - extra_data=body.get("extra_data"), + result = await agent_manager.invoke_agent( + source_session_id=session_id, message=message, + parent_session_id=body.get("parent_session_id", "") or None, + dashboard_id=body.get("dashboard_id", "") or None, ) - if result.get("success"): - from backend.apps.analytics.collector import record as _analytics - _analytics("subscription.connected", {"provider": provider}) - return result + return JSONResponse(result) + except ValueError as e: + return JSONResponse({"error": str(e)}, status_code=404) except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) - - -@agents.router.post("/subscriptions/exchange") -async def subscriptions_exchange(body: dict): - """Exchange OAuth code for tokens via 9Router.""" - from backend.apps.nine_router import exchange_oauth - provider = body.get("provider", "") - code = body.get("code", "") - redirect_uri = body.get("redirect_uri", "") - code_verifier = body.get("code_verifier", "") - state = body.get("state", "") - - if not provider or not code: - raise HTTPException(status_code=400, detail="provider and code required") - - try: - result = await exchange_oauth(provider, code, redirect_uri, code_verifier, state) - if result.get("success"): - from backend.apps.analytics.collector import record as _analytics - _analytics("subscription.connected", {"provider": provider}) - return result - except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) - - -@agents.router.get("/subscriptions/models") -async def subscriptions_models(): - """List all models available through connected subscriptions.""" - from backend.apps.nine_router import is_running, get_models - if not is_running(): - return {"models": []} - models = await get_models() - return {"models": models} - - -@agents.router.post("/subscriptions/disconnect") -async def subscriptions_disconnect(body: dict): - """Disconnect a subscription provider via 9Router.""" - import httpx - provider = body.get("provider", "") - if not provider: - raise HTTPException(status_code=400, detail="provider required") - - try: - from backend.apps.nine_router import NINE_ROUTER_API, get_providers - providers_data = await get_providers() - connections = providers_data.get("connections", []) if isinstance(providers_data, dict) else [] - conn = next((c for c in connections if c.get("provider") == provider), None) - if conn and conn.get("id"): - async with httpx.AsyncClient(timeout=10.0) as client: - await client.delete(f"{NINE_ROUTER_API}/providers/{conn['id']}") - from backend.apps.analytics.collector import record as _analytics - _analytics("subscription.disconnected", {"provider": provider}) - return {"ok": True} - return {"ok": False, "error": "Connection not found"} - except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) + return JSONResponse({"error": str(e)}, status_code=500) diff --git a/backend/apps/agents/approval.py b/backend/apps/agents/approval.py new file mode 100644 index 00000000..473cc842 --- /dev/null +++ b/backend/apps/agents/approval.py @@ -0,0 +1,79 @@ +"""Unified HITL (human-in-the-loop) approval flow. + +Used by both the main agent loop and browser sub-agents. +""" + +from __future__ import annotations + +import asyncio +import logging +from datetime import datetime +from uuid import uuid4 + +from backend.apps.agents.models import AgentSession, ApprovalRequest +from backend.apps.agents.ws_manager import ws_manager + +logger = logging.getLogger(__name__) + + +async def request_approval( + session: AgentSession, + tool_name: str, + tool_input: dict, + timeout: float | None = None, + track_analytics: bool = True, +) -> dict: + """Unified HITL approval flow. + + Creates an ApprovalRequest, sends it via WebSocket, waits for the user's + decision, cleans up, and returns the decision dict. + + Returns: {"behavior": "allow"|"deny", "message": ..., "updated_input": ...} + """ + safe_input = tool_input if isinstance(tool_input, dict) else {} + request_id = uuid4().hex + approval_req = ApprovalRequest( + id=request_id, session_id=session.id, + tool_name=tool_name, tool_input=safe_input, + ) + session.pending_approvals.append(approval_req) + session.status = "waiting_approval" + + if track_analytics: + from backend.apps.analytics.collector import record as _analytics + _analytics("approval.requested", { + "tool_name": tool_name, + "is_first_approval_in_session": len(session.pending_approvals) == 1, + "model": session.model, + }, session_id=session.id, dashboard_id=session.dashboard_id) + + await ws_manager.emit_status(session.id, "waiting_approval") + + if timeout is not None: + try: + decision = await asyncio.wait_for( + ws_manager.send_approval_request(session.id, request_id, tool_name, safe_input), + timeout=timeout, + ) + except asyncio.TimeoutError: + decision = {"behavior": "deny", "message": "Approval timed out"} + else: + decision = await ws_manager.send_approval_request( + session.id, request_id, tool_name, safe_input, + ) + + if track_analytics: + from backend.apps.analytics.collector import record as _analytics + latency_ms = int((datetime.now() - approval_req.created_at).total_seconds() * 1000) + _analytics("approval.resolved", { + "tool_name": tool_name, + "decision": decision.get("behavior", "unknown"), + "latency_ms": latency_ms, + "input_was_modified": decision.get("updated_input") is not None, + "model": session.model, + }, session_id=session.id, dashboard_id=session.dashboard_id) + + session.pending_approvals = [a for a in session.pending_approvals if a.id != request_id] + session.status = "running" + await ws_manager.emit_status(session.id, "running") + return decision diff --git a/backend/apps/agents/browser/executor.py b/backend/apps/agents/browser/executor.py index 519cd125..000584d5 100644 --- a/backend/apps/agents/browser/executor.py +++ b/backend/apps/agents/browser/executor.py @@ -2,13 +2,13 @@ from __future__ import annotations -import asyncio import json import logging from uuid import uuid4 -from backend.apps.agents.models import AgentSession, ApprovalRequest +from backend.apps.agents.models import AgentSession from backend.apps.agents.ws_manager import ws_manager +from backend.apps.agents.approval import request_approval from backend.apps.agents.browser.schemas import ACTION_MAP logger = logging.getLogger(__name__) @@ -46,26 +46,6 @@ def _format_tool_result(result: dict, tool_name: str) -> list[dict]: async def _request_browser_approval( session: AgentSession, tool_name: str, tool_input: dict, ) -> dict: - request_id = uuid4().hex - approval_req = ApprovalRequest( - id=request_id, session_id=session.id, - tool_name=tool_name, tool_input=tool_input, + return await request_approval( + session, tool_name, tool_input, timeout=300.0, track_analytics=False, ) - session.pending_approvals.append(approval_req) - session.status = "waiting_approval" - await ws_manager.send_to_session(session.id, "agent:status", { - "session_id": session.id, "status": "waiting_approval", - }) - try: - decision = await asyncio.wait_for( - ws_manager.send_approval_request(session.id, request_id, tool_name, tool_input), - timeout=300.0, - ) - except asyncio.TimeoutError: - decision = {"behavior": "deny", "message": "Approval timed out"} - session.pending_approvals = [a for a in session.pending_approvals if a.id != request_id] - session.status = "running" - await ws_manager.send_to_session(session.id, "agent:status", { - "session_id": session.id, "status": "running", - }) - return decision diff --git a/backend/apps/agents/browser/runner.py b/backend/apps/agents/browser/runner.py index 57942e42..7dfaaaf8 100644 --- a/backend/apps/agents/browser/runner.py +++ b/backend/apps/agents/browser/runner.py @@ -42,10 +42,7 @@ async def run_browser_agent( session._cancel_event = cancel_event agent_manager.sessions[session_id] = session - await ws_manager.send_to_session(session_id, "agent:status", { - "session_id": session_id, "status": "running", - "session": session.model_dump(mode="json"), - }) + await ws_manager.emit_status(session_id, "running", session) if initial_url: nav_result = await execute_browser_tool("BrowserNavigate", {"url": initial_url}, browser_id, tab_id) @@ -62,9 +59,7 @@ async def run_browser_agent( user_msg = Message(role="user", content=task) session.messages.append(user_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": user_msg.model_dump(mode="json"), - }) + await ws_manager.emit_message(session_id, user_msg) try: for turn in range(MAX_TURNS): @@ -88,15 +83,11 @@ async def run_browser_agent( if text_parts: asst_msg = Message(role="assistant", content="\n".join(text_parts)) session.messages.append(asst_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": asst_msg.model_dump(mode="json"), - }) + await ws_manager.emit_message(session_id, asst_msg) for tu in tool_uses: tool_msg = Message(role="tool_call", content={"id": tu.id, "tool": tu.name, "input": tu.input}) session.messages.append(tool_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, "message": tool_msg.model_dump(mode="json"), - }) + await ws_manager.emit_message(session_id, tool_msg) messages.append({"role": "assistant", "content": assistant_content}) if response.stop_reason != "tool_use": @@ -114,7 +105,7 @@ async def run_browser_agent( tool_results.append({"type": "tool_result", "tool_use_id": tu.id, "content": [{"type": "text", "text": denied_text}]}) result_msg = Message(role="tool_result", content={"text": denied_text, "tool_name": tu.name, "elapsed_ms": 0}) session.messages.append(result_msg) - await ws_manager.send_to_session(session_id, "agent:message", {"session_id": session_id, "message": result_msg.model_dump(mode="json")}) + await ws_manager.emit_message(session_id, result_msg) continue if policy == "ask": decision = await _request_browser_approval(session, tu.name, tu.input) @@ -123,7 +114,7 @@ async def run_browser_agent( tool_results.append({"type": "tool_result", "tool_use_id": tu.id, "content": [{"type": "text", "text": denied_text}]}) result_msg = Message(role="tool_result", content={"text": denied_text, "tool_name": tu.name, "elapsed_ms": 0}) session.messages.append(result_msg) - await ws_manager.send_to_session(session_id, "agent:message", {"session_id": session_id, "message": result_msg.model_dump(mode="json")}) + await ws_manager.emit_message(session_id, result_msg) continue start = time.time() @@ -137,7 +128,7 @@ async def run_browser_agent( result_text = result.get("text", result.get("error", "")) result_msg = Message(role="tool_result", content={"text": result_text, "tool_name": tu.name, "elapsed_ms": elapsed_ms}) session.messages.append(result_msg) - await ws_manager.send_to_session(session_id, "agent:message", {"session_id": session_id, "message": result_msg.model_dump(mode="json")}) + await ws_manager.emit_message(session_id, result_msg) messages.append({"role": "user", "content": tool_results}) if cancelled: @@ -145,7 +136,7 @@ async def run_browser_agent( if cancel_event.is_set(): session.status = "stopped" - await ws_manager.send_to_session(session_id, "agent:status", {"session_id": session_id, "status": "stopped", "session": session.model_dump(mode="json")}) + await ws_manager.emit_status(session_id, "stopped", session) return {"session_id": session_id, "browser_id": browser_id, "summary": "Agent was stopped.", "action_log": action_log, "final_screenshot": final_screenshot} summary_parts = text_parts if text_parts else ["Task completed."] @@ -160,7 +151,7 @@ async def run_browser_agent( pass session.status = "completed" - await ws_manager.send_to_session(session_id, "agent:status", {"session_id": session_id, "status": "completed", "session": session.model_dump(mode="json")}) + await ws_manager.emit_status(session_id, "completed", session) return {"session_id": session_id, "browser_id": browser_id, "summary": summary, "action_log": action_log, "final_screenshot": final_screenshot} except Exception as e: @@ -168,8 +159,8 @@ async def run_browser_agent( session.status = "error" error_msg = Message(role="system", content=f"Error: {str(e)}") session.messages.append(error_msg) - await ws_manager.send_to_session(session_id, "agent:message", {"session_id": session_id, "message": error_msg.model_dump(mode="json")}) - await ws_manager.send_to_session(session_id, "agent:status", {"session_id": session_id, "status": "error", "session": session.model_dump(mode="json")}) + await ws_manager.emit_message(session_id, error_msg) + await ws_manager.emit_status(session_id, "error", session) return {"session_id": session_id, "browser_id": browser_id, "summary": f"Error: {str(e)}", "action_log": action_log, "final_screenshot": None} diff --git a/backend/apps/agents/browser_agent_mcp_server.py b/backend/apps/agents/browser_agent_mcp_server.py index 0f7f4b6b..3ed7b353 100644 --- a/backend/apps/agents/browser_agent_mcp_server.py +++ b/backend/apps/agents/browser_agent_mcp_server.py @@ -21,7 +21,7 @@ except ImportError: HAS_PIL = False BACKEND_PORT = os.environ.get("OPENSWARM_PORT", "8324") -BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/browser-agent/run" +BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/agents/browser-agent/run" MODEL = os.environ.get("OPENSWARM_AGENT_MODEL", "sonnet") DASHBOARD_ID = os.environ.get("OPENSWARM_DASHBOARD_ID", "") PRE_SELECTED_BROWSER_IDS = os.environ.get("OPENSWARM_PRE_SELECTED_BROWSER_IDS", "") diff --git a/backend/apps/agents/browser_mcp_server.py b/backend/apps/agents/browser_mcp_server.py index 86f1ad5e..604e937a 100644 --- a/backend/apps/agents/browser_mcp_server.py +++ b/backend/apps/agents/browser_mcp_server.py @@ -22,7 +22,7 @@ except ImportError: HAS_PIL = False BACKEND_PORT = os.environ.get("OPENSWARM_PORT", "8324") -BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/browser/command" +BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/agents/browser/command" TAB_ID_PROP = { "type": "string", diff --git a/backend/apps/agents/invoke_agent_mcp_server.py b/backend/apps/agents/invoke_agent_mcp_server.py index cf614b03..68a2f1ce 100644 --- a/backend/apps/agents/invoke_agent_mcp_server.py +++ b/backend/apps/agents/invoke_agent_mcp_server.py @@ -14,7 +14,7 @@ import urllib.request import urllib.error BACKEND_PORT = os.environ.get("OPENSWARM_PORT", "8324") -BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/invoke-agent/run" +BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/agents/invoke-agent/run" PARENT_SESSION_ID = os.environ.get("OPENSWARM_PARENT_SESSION_ID", "") DASHBOARD_ID = os.environ.get("OPENSWARM_DASHBOARD_ID", "") diff --git a/backend/apps/agents/ws_manager.py b/backend/apps/agents/ws_manager.py index 24690628..abbfde21 100644 --- a/backend/apps/agents/ws_manager.py +++ b/backend/apps/agents/ws_manager.py @@ -123,4 +123,75 @@ class ConnectionManager: if future and not future.done(): future.set_result(result) + # ------------------------------------------------------------------ + # Typed event emitters + # ------------------------------------------------------------------ + + async def emit_status(self, session_id: str, status: str, session=None): + data: dict = {"session_id": session_id, "status": status} + if session is not None: + data["session"] = session.model_dump(mode="json") if hasattr(session, "model_dump") else session + await self.send_to_session(session_id, "agent:status", data) + + async def emit_message(self, session_id: str, message): + dumped = message.model_dump(mode="json") if hasattr(message, "model_dump") else message + await self.send_to_session(session_id, "agent:message", { + "session_id": session_id, "message": dumped, + }) + + async def emit_cost_update(self, session_id: str, cost_usd: float): + await self.send_to_session(session_id, "agent:cost_update", { + "session_id": session_id, "cost_usd": cost_usd, + }) + + async def emit_stream_start(self, session_id: str, message_id: str, role: str, tool_name: str = ""): + payload: dict = {"session_id": session_id, "message_id": message_id, "role": role} + if tool_name: + payload["tool_name"] = tool_name + await self.send_to_session(session_id, "agent:stream_start", payload) + + async def emit_stream_delta(self, session_id: str, message_id: str, delta: str): + await self.send_to_session(session_id, "agent:stream_delta", { + "session_id": session_id, "message_id": message_id, "delta": delta, + }) + + async def emit_stream_end(self, session_id: str, message_id: str): + await self.send_to_session(session_id, "agent:stream_end", { + "session_id": session_id, "message_id": message_id, + }) + + async def emit_branch_created(self, session_id: str, branch, active_branch_id: str): + dumped = branch.model_dump(mode="json") if hasattr(branch, "model_dump") else branch + await self.send_to_session(session_id, "agent:branch_created", { + "session_id": session_id, "branch": dumped, "active_branch_id": active_branch_id, + }) + + async def emit_branch_switched(self, session_id: str, active_branch_id: str): + await self.send_to_session(session_id, "agent:branch_switched", { + "session_id": session_id, "active_branch_id": active_branch_id, + }) + + async def emit_name_updated(self, session_id: str, name: str): + await self.send_to_session(session_id, "agent:name_updated", { + "session_id": session_id, "name": name, + }) + + async def emit_group_meta_updated( + self, session_id: str, group_id: str, name: str, svg: str, is_refined: bool, + ): + await self.send_to_session(session_id, "agent:group_meta_updated", { + "session_id": session_id, "group_id": group_id, + "name": name, "svg": svg, "is_refined": is_refined, + }) + + async def emit_closed(self, session_id: str, session): + await self.send_to_session(session_id, "agent:closed", { + "session_id": session_id, "status": session.status, + "name": session.name, "model": session.model, "mode": session.mode, + "created_at": session.created_at.isoformat() if session.created_at else None, + "closed_at": session.closed_at.isoformat() if session.closed_at else None, + "cost_usd": session.cost_usd, "dashboard_id": session.dashboard_id, + }) + + ws_manager = ConnectionManager() diff --git a/backend/apps/agents/ws_routes.py b/backend/apps/agents/ws_routes.py new file mode 100644 index 00000000..9752712c --- /dev/null +++ b/backend/apps/agents/ws_routes.py @@ -0,0 +1,60 @@ +"""WebSocket message dispatch logic. + +Extracted from main.py to keep it slim. The main.py WebSocket handlers +are thin wrappers that delegate here after JSON-parsing the message. +""" + +from __future__ import annotations + +import logging + +from backend.apps.agents.ws_manager import ws_manager + +logger = logging.getLogger(__name__) + + +async def handle_session_message(session_id: str, event: str, payload: dict): + """Dispatch an incoming WebSocket message for a session.""" + if event == "agent:send_message": + from backend.apps.agents.agent_manager import agent_manager + await agent_manager.send_message( + session_id, + payload.get("prompt", ""), + mode=payload.get("mode"), + model=payload.get("model"), + provider=payload.get("provider"), + images=payload.get("images"), + ) + elif event == "agent:approval_response": + from backend.apps.agents.agent_manager import agent_manager + agent_manager.handle_approval(payload.get("request_id"), { + "behavior": payload.get("behavior", "deny"), + "message": payload.get("message"), + "updated_input": payload.get("updated_input"), + }) + elif event == "agent:edit_message": + from backend.apps.agents.agent_manager import agent_manager + await agent_manager.edit_message( + session_id, + payload.get("message_id", ""), + payload.get("content", ""), + ) + elif event == "agent:stop": + from backend.apps.agents.agent_manager import agent_manager + await agent_manager.stop_agent(session_id) + + +async def handle_dashboard_message(event: str, payload: dict): + """Dispatch an incoming WebSocket message for the dashboard.""" + if event == "agent:approval_response": + from backend.apps.agents.agent_manager import agent_manager + agent_manager.handle_approval(payload.get("request_id"), { + "behavior": payload.get("behavior", "deny"), + "message": payload.get("message"), + "updated_input": payload.get("updated_input"), + }) + elif event == "browser:result": + ws_manager.resolve_browser_command( + payload.get("request_id", ""), + payload, + ) diff --git a/backend/apps/subscriptions/__init__.py b/backend/apps/subscriptions/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/apps/subscriptions/subscriptions.py b/backend/apps/subscriptions/subscriptions.py new file mode 100644 index 00000000..2be8a6e2 --- /dev/null +++ b/backend/apps/subscriptions/subscriptions.py @@ -0,0 +1,192 @@ +"""9Router / subscription management endpoints. + +Moved from agents.py and main.py to a dedicated sub-app. +""" + +from __future__ import annotations + +import logging +from contextlib import asynccontextmanager + +from fastapi import HTTPException, Request +from fastapi.responses import JSONResponse, HTMLResponse + +from backend.config.Apps import SubApp + +logger = logging.getLogger(__name__) + +_pending_oauth: dict[str, dict] = {} + + +@asynccontextmanager +async def subscriptions_lifespan(): + logger.info("Subscriptions sub-app starting") + yield + logger.info("Subscriptions sub-app shutting down") + + +subscriptions = SubApp("subscriptions", subscriptions_lifespan) + + +@subscriptions.router.get("/status") +async def subscriptions_status(): + """Check if 9Router is running and list connected providers.""" + from backend.apps.nine_router import is_running, get_providers, get_models + if not is_running(): + return {"running": False, "providers": [], "models": []} + providers = await get_providers() + models = await get_models() + return {"running": True, "providers": providers, "models": models} + + +@subscriptions.router.post("/connect") +async def subscriptions_connect(body: dict): + """Start OAuth flow for a subscription provider.""" + from backend.apps.nine_router import is_running, ensure_running, start_oauth + provider = body.get("provider", "") + if not provider: + raise HTTPException(status_code=400, detail="provider required") + + if not is_running(): + await ensure_running() + if not is_running(): + raise HTTPException(status_code=503, detail="9Router not available. Please install Node.js.") + + try: + result = await start_oauth(provider) + if result.get("flow") == "authorization_code" and result.get("state"): + _pending_oauth[result["state"]] = { + "provider": provider, + "code_verifier": result.get("code_verifier", ""), + "redirect_uri": result.get("redirect_uri", ""), + } + return result + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) + + +@subscriptions.router.post("/poll") +async def subscriptions_poll(body: dict): + """Poll for OAuth completion.""" + from backend.apps.nine_router import poll_oauth + provider = body.get("provider", "") + device_code = body.get("device_code", "") + if not provider or not device_code: + raise HTTPException(status_code=400, detail="provider and device_code required") + + try: + result = await poll_oauth( + provider, device_code, + code_verifier=body.get("code_verifier"), + extra_data=body.get("extra_data"), + ) + if result.get("success"): + from backend.apps.analytics.collector import record as _analytics + _analytics("subscription.connected", {"provider": provider}) + return result + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) + + +@subscriptions.router.post("/exchange") +async def subscriptions_exchange(body: dict): + """Exchange OAuth code for tokens via 9Router.""" + from backend.apps.nine_router import exchange_oauth + provider = body.get("provider", "") + code = body.get("code", "") + redirect_uri = body.get("redirect_uri", "") + code_verifier = body.get("code_verifier", "") + state = body.get("state", "") + + if not provider or not code: + raise HTTPException(status_code=400, detail="provider and code required") + + try: + result = await exchange_oauth(provider, code, redirect_uri, code_verifier, state) + if result.get("success"): + from backend.apps.analytics.collector import record as _analytics + _analytics("subscription.connected", {"provider": provider}) + return result + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) + + +@subscriptions.router.get("/models") +async def subscriptions_models(): + """List all models available through connected subscriptions.""" + from backend.apps.nine_router import is_running, get_models + if not is_running(): + return {"models": []} + models = await get_models() + return {"models": models} + + +@subscriptions.router.post("/disconnect") +async def subscriptions_disconnect(body: dict): + """Disconnect a subscription provider via 9Router.""" + import httpx + provider = body.get("provider", "") + if not provider: + raise HTTPException(status_code=400, detail="provider required") + + try: + from backend.apps.nine_router import NINE_ROUTER_API, get_providers + providers_data = await get_providers() + connections = providers_data.get("connections", []) if isinstance(providers_data, dict) else [] + conn = next((c for c in connections if c.get("provider") == provider), None) + if conn and conn.get("id"): + async with httpx.AsyncClient(timeout=10.0) as client: + await client.delete(f"{NINE_ROUTER_API}/providers/{conn['id']}") + from backend.apps.analytics.collector import record as _analytics + _analytics("subscription.disconnected", {"provider": provider}) + return {"ok": True} + return {"ok": False, "error": "Connection not found"} + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) + + +@subscriptions.router.get("/pending/{state}") +async def subscriptions_pending(state: str): + """Return pending OAuth data for a state param. Called by 9Router's callback page.""" + pending = _pending_oauth.get(state) + if not pending: + return JSONResponse({"error": "not found"}, status_code=404, + headers={"Access-Control-Allow-Origin": "*"}) + return JSONResponse({ + "provider": pending["provider"], + "code_verifier": pending["code_verifier"], + "redirect_uri": pending["redirect_uri"], + }, headers={"Access-Control-Allow-Origin": "*"}) + + +@subscriptions.router.get("/callback") +async def subscriptions_callback(request: Request): + """Catch OAuth redirect from provider, exchange code via 9Router, close window.""" + code = request.query_params.get("code", "") + state = request.query_params.get("state", "") + error = request.query_params.get("error", "") + + if error: + desc = request.query_params.get("error_description", error) + return HTMLResponse(f'

Authorization failed

{desc}

') + + pending = _pending_oauth.pop(state, None) + if not pending: + return HTMLResponse('

Session expired

Please try connecting again.

') + + from backend.apps.nine_router import exchange_oauth + try: + await exchange_oauth(pending["provider"], code, pending["redirect_uri"], pending["code_verifier"], state) + except Exception as e: + return HTMLResponse(f'

Connection failed

{e}

') + + return HTMLResponse( + '' + '
' + '
' + '

Connected!

' + '

You can close this window

' + '
' + '' + '' + ) diff --git a/backend/main.py b/backend/main.py index 2d1284c3..d34bdf30 100644 --- a/backend/main.py +++ b/backend/main.py @@ -1,18 +1,13 @@ import logging import os -from uuid import uuid4 logger = logging.getLogger(__name__) -from fastapi.responses import JSONResponse, HTMLResponse -from fastapi import Request - -# In-memory store for pending OAuth flows (state -> {provider, code_verifier, redirect_uri}) -_pending_oauth: dict[str, dict] = {} from backend.config.Apps import MainApp from backend.apps.health.health import health from backend.apps.agents.agents import agents from backend.apps.agents.ws_manager import ws_manager +from backend.apps.agents.ws_routes import handle_session_message, handle_dashboard_message from backend.apps.templates.templates import templates from backend.apps.skills.skills import skills from backend.apps.tools_lib.tools_lib import tools_lib @@ -23,11 +18,16 @@ from backend.apps.skill_registry.skill_registry import skill_registry from backend.apps.outputs.outputs import outputs from backend.apps.dashboards.dashboards import dashboards from backend.apps.analytics.analytics import analytics +from backend.apps.subscriptions.subscriptions import subscriptions from fastapi.middleware.cors import CORSMiddleware from fastapi import WebSocket, WebSocketDisconnect import json -main_app = MainApp([health, agents, templates, skills, tools_lib, modes, settings, mcp_registry, skill_registry, outputs, dashboards, analytics]) +main_app = MainApp([ + health, agents, templates, skills, tools_lib, modes, settings, + mcp_registry, skill_registry, outputs, dashboards, analytics, + subscriptions, +]) app = main_app.app app.add_middleware( @@ -38,6 +38,7 @@ app.add_middleware( allow_headers=["*"], ) + @app.websocket("/ws/agents/{session_id}") async def websocket_session(websocket: WebSocket, session_id: str): await ws_manager.connect_session(session_id, websocket) @@ -45,39 +46,11 @@ async def websocket_session(websocket: WebSocket, session_id: str): while True: data = await websocket.receive_text() msg = json.loads(data) - event = msg.get("event") - payload = msg.get("data", {}) - - if event == "agent:send_message": - from backend.apps.agents.agent_manager import agent_manager - await agent_manager.send_message( - session_id, - payload.get("prompt", ""), - mode=payload.get("mode"), - model=payload.get("model"), - provider=payload.get("provider"), - images=payload.get("images"), - ) - elif event == "agent:approval_response": - from backend.apps.agents.agent_manager import agent_manager - agent_manager.handle_approval(payload.get("request_id"), { - "behavior": payload.get("behavior", "deny"), - "message": payload.get("message"), - "updated_input": payload.get("updated_input"), - }) - elif event == "agent:edit_message": - from backend.apps.agents.agent_manager import agent_manager - await agent_manager.edit_message( - session_id, - payload.get("message_id", ""), - payload.get("content", ""), - ) - elif event == "agent:stop": - from backend.apps.agents.agent_manager import agent_manager - await agent_manager.stop_agent(session_id) + await handle_session_message(session_id, msg.get("event"), msg.get("data", {})) except WebSocketDisconnect: ws_manager.disconnect_session(session_id, websocket) + @app.websocket("/ws/dashboard") async def websocket_dashboard(websocket: WebSocket): await ws_manager.connect_global(websocket) @@ -85,148 +58,11 @@ async def websocket_dashboard(websocket: WebSocket): while True: data = await websocket.receive_text() msg = json.loads(data) - event = msg.get("event") - payload = msg.get("data", {}) - - if event == "agent:approval_response": - from backend.apps.agents.agent_manager import agent_manager - agent_manager.handle_approval(payload.get("request_id"), { - "behavior": payload.get("behavior", "deny"), - "message": payload.get("message"), - "updated_input": payload.get("updated_input"), - }) - elif event == "browser:result": - ws_manager.resolve_browser_command( - payload.get("request_id", ""), - payload, - ) + await handle_dashboard_message(msg.get("event"), msg.get("data", {})) except WebSocketDisconnect: ws_manager.disconnect_global(websocket) -@app.post("/api/browser/command") -async def browser_command(request: Request): - """HTTP endpoint called by the browser MCP server subprocess. - Proxies commands to the frontend via WebSocket and waits for results.""" - body = await request.json() - action = body.get("action", "") - browser_id = body.get("browser_id", "") - tab_id = body.get("tab_id", "") - params = body.get("params", {}) - - if not action or not browser_id: - return JSONResponse({"error": "action and browser_id are required"}, status_code=400) - - request_id = uuid4().hex - result = await ws_manager.send_browser_command(request_id, action, browser_id, params, tab_id=tab_id) - return JSONResponse(result) - - -@app.get("/api/subscriptions/pending/{state}") -async def subscriptions_pending(state: str): - """Return pending OAuth data for a state param. Called by 9Router's callback page.""" - pending = _pending_oauth.get(state) - if not pending: - return JSONResponse({"error": "not found"}, status_code=404, - headers={"Access-Control-Allow-Origin": "*"}) - return JSONResponse({ - "provider": pending["provider"], - "code_verifier": pending["code_verifier"], - "redirect_uri": pending["redirect_uri"], - }, headers={"Access-Control-Allow-Origin": "*"}) - - -@app.get("/api/subscriptions/callback") -async def subscriptions_callback(request: Request): - """Catch OAuth redirect from provider, exchange code via 9Router, close window.""" - code = request.query_params.get("code", "") - state = request.query_params.get("state", "") - error = request.query_params.get("error", "") - - if error: - desc = request.query_params.get("error_description", error) - return HTMLResponse(f'

Authorization failed

{desc}

') - - pending = _pending_oauth.pop(state, None) - if not pending: - return HTMLResponse('

Session expired

Please try connecting again.

') - - from backend.apps.nine_router import exchange_oauth - try: - await exchange_oauth(pending["provider"], code, pending["redirect_uri"], pending["code_verifier"], state) - except Exception as e: - return HTMLResponse(f'

Connection failed

{e}

') - - return HTMLResponse( - '' - '
' - '
' - '

Connected!

' - '

You can close this window

' - '
' - '' - '' - ) - - -@app.post("/api/browser-agent/run") -async def browser_agent_run(request: Request): - """Run one or more browser sub-agents in parallel. - Called by the browser_agent_mcp_server stdio subprocess.""" - from backend.apps.settings.settings import load_settings - from backend.apps.agents.browser_agent import run_browser_agents - - body = await request.json() - tasks = body.get("tasks", []) - model = body.get("model", "sonnet") - dashboard_id = body.get("dashboard_id", "") - pre_selected_browser_ids = body.get("pre_selected_browser_ids", []) - parent_session_id = body.get("parent_session_id", "") - - if not tasks: - return JSONResponse({"error": "tasks array is required"}, status_code=400) - - results = await run_browser_agents( - tasks=tasks, - model=model, - dashboard_id=dashboard_id or None, - pre_selected_browser_ids=pre_selected_browser_ids, - parent_session_id=parent_session_id or None, - ) - return JSONResponse({"results": results}) - - -@app.post("/api/invoke-agent/run") -async def invoke_agent_run(request: Request): - """Fork an existing agent session and send it a new message. - Called by the invoke_agent_mcp_server stdio subprocess.""" - body = await request.json() - session_id = body.get("session_id", "") - message = body.get("message", "") - parent_session_id = body.get("parent_session_id", "") - dashboard_id = body.get("dashboard_id", "") - - if not session_id: - return JSONResponse({"error": "session_id is required"}, status_code=400) - if not message: - return JSONResponse({"error": "message is required"}, status_code=400) - - try: - from backend.apps.agents.agent_manager import agent_manager - result = await agent_manager.invoke_agent( - source_session_id=session_id, - message=message, - parent_session_id=parent_session_id or None, - dashboard_id=dashboard_id or None, - ) - return JSONResponse(result) - except ValueError as e: - return JSONResponse({"error": str(e)}, status_code=404) - except Exception as e: - logger.exception("invoke_agent_run failed") - return JSONResponse({"error": str(e)}, status_code=500) - - if __name__ == "__main__": import argparse import uvicorn diff --git a/frontend/src/app/components/OnboardingModal.tsx b/frontend/src/app/components/OnboardingModal.tsx index baac74b7..ec16d03d 100644 --- a/frontend/src/app/components/OnboardingModal.tsx +++ b/frontend/src/app/components/OnboardingModal.tsx @@ -40,7 +40,7 @@ const OnboardingModal: React.FC = () => { let attempts = 0; const maxAttempts = 15; // 30 seconds const check = () => { - fetch(`${API_BASE}/agents/subscriptions/status`) + fetch(`${API_BASE}/subscriptions/status`) .then((r) => r.json()) .then((data) => { if (data.running) { @@ -106,7 +106,7 @@ const OnboardingModal: React.FC = () => { await new Promise(r => setTimeout(r, 1000)); try { - const r = await fetch(`${API_BASE}/agents/subscriptions/connect`, { + const r = await fetch(`${API_BASE}/subscriptions/connect`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ provider: providerId }), @@ -122,7 +122,7 @@ const OnboardingModal: React.FC = () => { const timer = setInterval(async () => { try { - const pr = await fetch(`${API_BASE}/agents/subscriptions/poll`, { + const pr = await fetch(`${API_BASE}/subscriptions/poll`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ @@ -149,7 +149,7 @@ const OnboardingModal: React.FC = () => { // Poll status as primary detection (works in Electron where postMessage may not) const statusPoller = setInterval(async () => { try { - const sr = await fetch(`${API_BASE}/agents/subscriptions/status`); + const sr = await fetch(`${API_BASE}/subscriptions/status`); const sd = await sr.json(); const connections = sd.providers?.connections || []; if (connections.some((p: any) => p.provider === providerId && p.isActive)) { @@ -175,7 +175,7 @@ const OnboardingModal: React.FC = () => { if (pollTimerRef.current) { clearInterval(pollTimerRef.current); pollTimerRef.current = null; } if (popup && !popup.closed) popup.close(); try { - await fetch(`${API_BASE}/agents/subscriptions/exchange`, { + await fetch(`${API_BASE}/subscriptions/exchange`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ diff --git a/frontend/src/app/pages/Settings/Settings.tsx b/frontend/src/app/pages/Settings/Settings.tsx index f4bfb0f4..9079c0c1 100644 --- a/frontend/src/app/pages/Settings/Settings.tsx +++ b/frontend/src/app/pages/Settings/Settings.tsx @@ -237,7 +237,7 @@ const SubscriptionCards: React.FC = () => { const [pollTimer, setPollTimer] = useState(null); const fetchStatus = () => { - fetch(`${API_BASE}/agents/subscriptions/status`) + fetch(`${API_BASE}/subscriptions/status`) .then(r => r.json()) .then(setStatus) .catch(() => setStatus({ running: false, providers: [], models: [] })); @@ -261,7 +261,7 @@ const SubscriptionCards: React.FC = () => { await new Promise(r => setTimeout(r, 500)); try { - const r = await fetch(`${API_BASE}/agents/subscriptions/connect`, { + const r = await fetch(`${API_BASE}/subscriptions/connect`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ provider: providerId }), }); @@ -275,7 +275,7 @@ const SubscriptionCards: React.FC = () => { const timer = setInterval(async () => { try { - const pr = await fetch(`${API_BASE}/agents/subscriptions/poll`, { + const pr = await fetch(`${API_BASE}/subscriptions/poll`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ provider: providerId, device_code: data.device_code, code_verifier: data.code_verifier, extra_data: data.extra_data }), }); @@ -298,7 +298,7 @@ const SubscriptionCards: React.FC = () => { // Status polling as primary detection const statusPoller = setInterval(async () => { try { - const sr = await fetch(`${API_BASE}/agents/subscriptions/status`); + const sr = await fetch(`${API_BASE}/subscriptions/status`); const sd = await sr.json(); const connections = sd.providers?.connections || []; if (connections.some((p: any) => p.provider === providerId && p.isActive)) { @@ -322,7 +322,7 @@ const SubscriptionCards: React.FC = () => { setPollTimer(null); if (popup && !popup.closed) popup.close(); try { - await fetch(`${API_BASE}/agents/subscriptions/exchange`, { + await fetch(`${API_BASE}/subscriptions/exchange`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ provider: providerId, code: callbackData.code, @@ -354,7 +354,7 @@ const SubscriptionCards: React.FC = () => { const handleDisconnect = async (providerId: string) => { setDisconnecting(providerId); try { - await fetch(`${API_BASE}/agents/subscriptions/disconnect`, { + await fetch(`${API_BASE}/subscriptions/disconnect`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ provider: providerId }),