mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-11 12:17:45 +02:00
[Haik]: Agentic refactor 3. DRY Up Cross-Cutting Patterns
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"mcpServers": {
|
||||
"mcp-docs-server": {
|
||||
"command": "npx",
|
||||
"args": [
|
||||
"-y",
|
||||
"@assistant-ui/mcp-docs-server"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
+128
-600
@@ -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"<app_builder_reference>\n{VIEW_BUILDER_SKILL}\n</app_builder_reference>"
|
||||
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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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"<app_builder_reference>\n{VIEW_BUILDER_SKILL}\n</app_builder_reference>"
|
||||
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
|
||||
+48
-118
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
|
||||
@@ -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", "")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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", "")
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Authorization failed</h2><p style="color:#888">{desc}</p></div></body></html>')
|
||||
|
||||
pending = _pending_oauth.pop(state, None)
|
||||
if not pending:
|
||||
return HTMLResponse('<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Session expired</h2><p style="color:#888">Please try connecting again.</p></div></body></html>')
|
||||
|
||||
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'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Connection failed</h2><p style="color:#888">{e}</p></div></body></html>')
|
||||
|
||||
return HTMLResponse(
|
||||
'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif">'
|
||||
'<div style="text-align:center">'
|
||||
'<div style="width:64px;height:64px;border-radius:50%;background:#22c55e20;display:flex;align-items:center;justify-content:center;margin:0 auto 16px;font-size:32px">✓</div>'
|
||||
'<h2 style="margin:0 0 8px">Connected!</h2>'
|
||||
'<p style="color:#888;margin:0">You can close this window</p>'
|
||||
'</div>'
|
||||
'<script>setTimeout(()=>window.close(),1500)</script>'
|
||||
'</body></html>'
|
||||
)
|
||||
+11
-175
@@ -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'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Authorization failed</h2><p style="color:#888">{desc}</p></div></body></html>')
|
||||
|
||||
pending = _pending_oauth.pop(state, None)
|
||||
if not pending:
|
||||
return HTMLResponse('<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Session expired</h2><p style="color:#888">Please try connecting again.</p></div></body></html>')
|
||||
|
||||
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'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Connection failed</h2><p style="color:#888">{e}</p></div></body></html>')
|
||||
|
||||
return HTMLResponse(
|
||||
'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif">'
|
||||
'<div style="text-align:center">'
|
||||
'<div style="width:64px;height:64px;border-radius:50%;background:#22c55e20;display:flex;align-items:center;justify-content:center;margin:0 auto 16px;font-size:32px">✓</div>'
|
||||
'<h2 style="margin:0 0 8px">Connected!</h2>'
|
||||
'<p style="color:#888;margin:0">You can close this window</p>'
|
||||
'</div>'
|
||||
'<script>setTimeout(()=>window.close(),1500)</script>'
|
||||
'</body></html>'
|
||||
)
|
||||
|
||||
|
||||
@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
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -237,7 +237,7 @@ const SubscriptionCards: React.FC = () => {
|
||||
const [pollTimer, setPollTimer] = useState<any>(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 }),
|
||||
|
||||
Reference in New Issue
Block a user