From 5536c0e60ef0b3a8bd7f778f7a7037e612a1f180 Mon Sep 17 00:00:00 2001 From: haikdc Date: Sun, 5 Apr 2026 14:24:16 -0700 Subject: [PATCH] [Haik]: tools subapp done, some touch ups may be done tho --- .../ResolvedModeConfig/ResolvedModeConfig.py | 7 +- .../compose_system_prompt.py | 6 +- .../agents/agent_utils/build_agent_toolkit.py | 52 ++++ backend/apps/agents/agents.py | 29 +- .../dashboards/generate_dashboard_name.py | 9 +- backend/apps/tools/builtin_tools.py | 50 ++++ .../tools/discover_tools/DiscoveryError.py | 8 + .../tools/discover_tools/discover_tools.py | 42 +++ .../utils/discover_mcp_tools_http.py | 70 +++++ .../utils/discover_mcp_tools_sse.py | 26 ++ .../utils/discover_mcp_tools_stdio.py | 97 +++++++ backend/apps/tools/oauth/oauth.py | 265 ++++++++++++++++++ backend/apps/tools/oauth/oauth_providers.py | 181 ++++++++++++ .../apps/tools/shared_utils/ToolDefinition.py | 23 ++ backend/apps/tools/shared_utils/mcp_config.py | 69 +++++ .../helpers/build_http_sse_tool.py | 22 ++ .../helpers/build_stdio_tool.py | 49 ++++ .../helpers/inject_credentials.py | 67 +++++ .../tool_definition_to_mcp_tool.py | 52 ++++ backend/apps/tools/tools.py | 255 +++++++++++++++++ backend/core/tools/shared_structs/MCP_Tool.py | 4 +- backend/main.py | 3 +- 22 files changed, 1366 insertions(+), 20 deletions(-) create mode 100644 backend/apps/agents/agent_utils/build_agent_toolkit.py create mode 100644 backend/apps/tools/builtin_tools.py create mode 100644 backend/apps/tools/discover_tools/DiscoveryError.py create mode 100644 backend/apps/tools/discover_tools/discover_tools.py create mode 100644 backend/apps/tools/discover_tools/utils/discover_mcp_tools_http.py create mode 100644 backend/apps/tools/discover_tools/utils/discover_mcp_tools_sse.py create mode 100644 backend/apps/tools/discover_tools/utils/discover_mcp_tools_stdio.py create mode 100644 backend/apps/tools/oauth/oauth.py create mode 100644 backend/apps/tools/oauth/oauth_providers.py create mode 100644 backend/apps/tools/shared_utils/ToolDefinition.py create mode 100644 backend/apps/tools/shared_utils/mcp_config.py create mode 100644 backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_http_sse_tool.py create mode 100644 backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_stdio_tool.py create mode 100644 backend/apps/tools/tool_definition_to_mcp_tool/helpers/inject_credentials.py create mode 100644 backend/apps/tools/tool_definition_to_mcp_tool/tool_definition_to_mcp_tool.py create mode 100644 backend/apps/tools/tools.py diff --git a/backend/apps/agents/ResolvedModeConfig/ResolvedModeConfig.py b/backend/apps/agents/ResolvedModeConfig/ResolvedModeConfig.py index 30c675c4..b6543019 100644 --- a/backend/apps/agents/ResolvedModeConfig/ResolvedModeConfig.py +++ b/backend/apps/agents/ResolvedModeConfig/ResolvedModeConfig.py @@ -1,8 +1,9 @@ -from typing import Optional, Tuple, List +from typing import Optional, List from backend.apps.modes.modes import get_mode_by_id from backend.apps.modes.Mode import Mode from backend.apps.settings.settings import load_settings +from backend.apps.settings.AppSettings import AppSettings from backend.apps.agents.ResolvedModeConfig.compose_system_prompt import compose_system_prompt from backend.core.tools.shared_structs.Toolkit import Toolkit from typeguard import typechecked @@ -19,11 +20,11 @@ class ResolvedModeConfig(BaseModel): async def create( cls, mode_id: str, - session_prompt: Optional[str], toolkit: Toolkit, + session_prompt: Optional[str] = None, ) -> "ResolvedModeConfig": - settings = load_settings() + settings: AppSettings = load_settings() mode_def: Optional[Mode] = await get_mode_by_id(mode_id) system_prompt: Optional[str] = compose_system_prompt( diff --git a/backend/apps/agents/ResolvedModeConfig/compose_system_prompt.py b/backend/apps/agents/ResolvedModeConfig/compose_system_prompt.py index 32548bf3..4ce4e118 100644 --- a/backend/apps/agents/ResolvedModeConfig/compose_system_prompt.py +++ b/backend/apps/agents/ResolvedModeConfig/compose_system_prompt.py @@ -12,8 +12,7 @@ def compose_system_prompt( global_default: Optional[str] = None, mode_prompt: Optional[str] = None, session_prompt: Optional[str] = None, - connected_tools_ctx: Optional[str] = None, - browser_ctx: Optional[str] = None, + browser_context: Optional[str] = None, ) -> Optional[str]: """Layer multiple prompt sources into one system prompt. @@ -24,7 +23,6 @@ def compose_system_prompt( global_default, mode_prompt, session_prompt, - connected_tools_ctx, - browser_ctx, + browser_context, ) if p] return "\n\n".join(parts) if parts else None \ No newline at end of file diff --git a/backend/apps/agents/agent_utils/build_agent_toolkit.py b/backend/apps/agents/agent_utils/build_agent_toolkit.py new file mode 100644 index 00000000..6dbd2ca8 --- /dev/null +++ b/backend/apps/agents/agent_utils/build_agent_toolkit.py @@ -0,0 +1,52 @@ +from backend.core.Agent.Agent import Agent +from backend.core.tools.shared_structs.Toolkit import Toolkit +from backend.core.tools.make_builtin_toolkit.make_builtin_toolkit import make_builtin_toolkit +from backend.apps.agents.COMMS_MANAGER.COMMS_MANAGER import CommsManager +from backend.apps.tools.tools import load_user_toolkit, load_builtin_permissions +from typing import Dict, Optional +from typeguard import typechecked +from backend.core.tools.shared_structs.TOOL_PERMISSIONS import TOOL_PERMISSIONS + + +@typechecked +def p_apply_builtin_permission_overrides(toolkit: Toolkit, permissions: dict[str, TOOL_PERMISSIONS]) -> None: + """Walk the toolkit tree and apply user-saved builtin permission overrides.""" + if toolkit.tools is not None: + for tool in toolkit.tools: + sdk_name: str = tool.to_sdk_args() + if sdk_name in permissions: + perm: TOOL_PERMISSIONS = permissions[sdk_name] + tool.permission = perm + if toolkit.nested_toolkits is not None: + for nested in toolkit.nested_toolkits: + p_apply_builtin_permission_overrides(nested, permissions) + + +@typechecked +def build_agent_toolkit(agent: Agent, sessions: Dict[str, Agent], comms_manager: CommsManager) -> Toolkit: + """Build the full toolkit tree: builtin tools + user-installed MCP tools. + + Also applies saved builtin permission overrides. + """ + builtin_toolkit: Toolkit = make_builtin_toolkit(agent, sessions, comms_manager.send_browser_command) + user_toolkit: Optional[Toolkit] = load_user_toolkit() + + full_toolkit: Optional[Toolkit] = None + if user_toolkit is not None: + full_toolkit = Toolkit( + name="root", + description="All available tools", + nested_toolkits=[builtin_toolkit, user_toolkit], + ) + else: + full_toolkit = builtin_toolkit + assert full_toolkit is not None, "Full toolkit should never be None" + + builtin_permissions: Optional[dict[str, TOOL_PERMISSIONS]] = load_builtin_permissions() + if builtin_permissions is not None: + p_apply_builtin_permission_overrides( + toolkit=full_toolkit, + permissions=builtin_permissions, + ) + + return full_toolkit \ No newline at end of file diff --git a/backend/apps/agents/agents.py b/backend/apps/agents/agents.py index 74919609..06a0c597 100644 --- a/backend/apps/agents/agents.py +++ b/backend/apps/agents/agents.py @@ -23,7 +23,6 @@ from backend.core.Agent.Agent import Agent from backend.core.db.PydanticStore import PydanticStore from backend.core.shared_structs.agent.Message.Message import UserMessage from backend.core.events.events import AgentStatusEvent, AgentClosedEvent, BranchSwitchedEvent -from backend.core.tools.make_builtin_toolkit.make_builtin_toolkit import make_builtin_toolkit from backend.apps.agents.agent_utils.create_sdk_hooks import create_sdk_hooks from backend.apps.agents.agent_utils.build_search_text import build_search_text from backend.apps.agents.ResolvedModeConfig.ResolvedModeConfig import ResolvedModeConfig @@ -34,6 +33,7 @@ from backend.ports import NINE_ROUTER_PORT from claude_agent_sdk import ClaudeAgentOptions from claude_agent_sdk.types import HookMatcher, McpServerConfig from backend.core.tools.shared_structs.Toolkit import Toolkit +from backend.apps.agents.agent_utils.build_agent_toolkit import build_agent_toolkit AGENT_STORE: PydanticStore[Agent] = PydanticStore[Agent]( model_cls=Agent, @@ -62,8 +62,11 @@ async def agents_lifespan(): try: stored.status = "stopped" stored.on_event = COMMS_MANAGER.make_session_emitter(stored.session_id) - toolkit: Toolkit = make_builtin_toolkit(stored, SESSIONS, COMMS_MANAGER.send_browser_command) - stored.toolkit = toolkit + stored.toolkit = build_agent_toolkit( + agent=stored, + sessions=SESSIONS, + comms_manager=COMMS_MANAGER, + ) SESSIONS[stored.session_id] = stored except Exception as e: print(f"[agents lifespan] Skipping corrupt session {stored.session_id}: {e}") @@ -111,7 +114,11 @@ async def launch(body: LaunchBody) -> dict: agent.on_event = COMMS_MANAGER.make_session_emitter(agent.session_id) SESSIONS[agent.session_id] = agent - toolkit: Toolkit = make_builtin_toolkit(agent, SESSIONS, COMMS_MANAGER.send_browser_command) + toolkit: Toolkit = build_agent_toolkit( + agent=agent, + sessions=SESSIONS, + comms_manager=COMMS_MANAGER, + ) agent.toolkit = toolkit mcp_servers: Dict[str, McpServerConfig] = toolkit.collect_mcp_servers() @@ -316,8 +323,11 @@ async def resume_session(session_id: str) -> dict: raise HTTPException(status_code=404, detail="Session not found in history") agent.status = "stopped" agent.on_event = COMMS_MANAGER.make_session_emitter(agent.session_id) - toolkit: Toolkit = make_builtin_toolkit(agent, SESSIONS, COMMS_MANAGER.send_browser_command) - agent.toolkit = toolkit + agent.toolkit = build_agent_toolkit( + agent=agent, + sessions=SESSIONS, + comms_manager=COMMS_MANAGER, + ) SESSIONS[agent.session_id] = agent AGENT_STORE.delete(session_id) await agent.emit(AgentStatusEvent( @@ -343,8 +353,11 @@ async def duplicate_session(session_id: str, body: dict = {}) -> dict: clone.pending_approvals = [] clone.sub_agents = [] clone.on_event = COMMS_MANAGER.make_session_emitter(clone.session_id) - toolkit: Toolkit = make_builtin_toolkit(clone, SESSIONS, COMMS_MANAGER.send_browser_command) - clone.toolkit = toolkit + clone.toolkit = build_agent_toolkit( + agent=clone, + sessions=SESSIONS, + comms_manager=COMMS_MANAGER, + ) SESSIONS[clone.session_id] = clone await clone.emit(AgentStatusEvent( session_id=clone.session_id, status=clone.status, diff --git a/backend/apps/dashboards/generate_dashboard_name.py b/backend/apps/dashboards/generate_dashboard_name.py index 5e8e012f..d4e7f859 100644 --- a/backend/apps/dashboards/generate_dashboard_name.py +++ b/backend/apps/dashboards/generate_dashboard_name.py @@ -2,8 +2,8 @@ from typing import List, Optional from backend.core.llm.quick_llm_call import quick_llm_call from typeguard import typechecked -SINGLE_AGENT_SYSTEM_PROMPT = "Generate a concise 2-5 word workspace name for a project based on this task. Return only the name with no markdown formatting, nothing else." -MULTI_AGENT_SYSTEM_PROMPT = "Generate a concise 2-5 word workspace name that captures the overall theme of these tasks. Return only the name with no markdown formatting, nothing else." +SINGLE_AGENT_SYSTEM_PROMPT: str = "Generate a concise 2-5 word workspace name for a project based on this task. Return only the name with no markdown formatting, nothing else." +MULTI_AGENT_SYSTEM_PROMPT: str = "Generate a concise 2-5 word workspace name that captures the overall theme of these tasks. Return only the name with no markdown formatting, nothing else." @typechecked async def generate_dashboard_name( @@ -12,12 +12,17 @@ async def generate_dashboard_name( nine_router_port: Optional[int] = None, ) -> str: """Given user prompts from a dashboard's sessions, generate a short name via LLM.""" + system_prompt: Optional[str] = None + user_prompt: Optional[str] = None if len(prompts) == 1: system_prompt = SINGLE_AGENT_SYSTEM_PROMPT user_prompt = prompts[0] else: system_prompt = MULTI_AGENT_SYSTEM_PROMPT user_prompt = "\n".join(f"- {p}" for p in prompts) + + assert system_prompt is not None, "System prompt cannot be None" + assert user_prompt is not None, "User prompt cannot be None" result = await quick_llm_call( api_key=api_key, diff --git a/backend/apps/tools/builtin_tools.py b/backend/apps/tools/builtin_tools.py new file mode 100644 index 00000000..85bb58e2 --- /dev/null +++ b/backend/apps/tools/builtin_tools.py @@ -0,0 +1,50 @@ +"""Flat list of builtin tool metadata for the /builtin endpoint. + +Derived from the existing pre_existing_toolkits constants so there is +a single source of truth for tool definitions. +""" + +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.basic_toolkits.FILESYSTEM_TOOLKIT import FILESYSTEM_TOOLKIT +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.basic_toolkits.SEARCH_TOOLKIT import SEARCH_TOOLKIT +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.basic_toolkits.SYSTEM_TOOLKIT import SYSTEM_TOOLKIT +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.meta_toolkits.INTERACTION_TOOLKIT import INTERACTION_TOOLKIT +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.meta_toolkits.PLANNING_TOOLKIT import PLANNING_TOOLKIT +from backend.core.tools.make_builtin_toolkit.pre_existing_toolkits.meta_toolkits.SCHEDULING_TOOLKIT import SCHEDULING_TOOLKIT +from backend.core.tools.shared_structs.Toolkit import Toolkit + + +def _collect_tools(toolkit: Toolkit, category: str) -> list[dict]: + if toolkit.tools is None: + return [] + return [ + { + "name": t.name, + "description": t.description or "", + "category": category, + "deferred": t.deferred, + } + for t in toolkit.tools + ] + + +BUILTIN_TOOLS: list[dict] = [ + *_collect_tools(FILESYSTEM_TOOLKIT, "filesystem"), + *_collect_tools(SYSTEM_TOOLKIT, "system"), + *_collect_tools(SEARCH_TOOLKIT, "search"), + *_collect_tools(INTERACTION_TOOLKIT, "interaction"), + *_collect_tools(PLANNING_TOOLKIT, "planning"), + *_collect_tools(SCHEDULING_TOOLKIT, "scheduling"), + {"name": "CreateAgent", "description": "Spawn a sub-agent to handle a complex subtask", "category": "agents", "deferred": False}, + {"name": "InvokeAgent", "description": "Invoke a copy of an existing agent with a new message", "category": "agents", "deferred": False}, + {"name": "CreateBrowserAgent", "description": "Create a new browser and run a task on it", "category": "browser_delegation", "deferred": False}, + {"name": "BrowserAgent", "description": "Delegate a browser task to an existing browser agent", "category": "browser_delegation", "deferred": False}, + {"name": "BrowserScreenshot", "description": "Capture a screenshot of the browser page", "category": "browser_action", "deferred": False}, + {"name": "BrowserNavigate", "description": "Navigate the browser to a URL", "category": "browser_action", "deferred": False}, + {"name": "BrowserClick", "description": "Click an element by CSS selector", "category": "browser_action", "deferred": False}, + {"name": "BrowserType", "description": "Type text into an input element", "category": "browser_action", "deferred": False}, + {"name": "BrowserEvaluate", "description": "Execute JavaScript in the browser", "category": "browser_action", "deferred": False}, + {"name": "BrowserGetText", "description": "Get visible text content of the page", "category": "browser_action", "deferred": False}, + {"name": "BrowserGetElements", "description": "List interactive elements with CSS selectors", "category": "browser_action", "deferred": False}, + {"name": "BrowserScroll", "description": "Scroll the page up or down", "category": "browser_action", "deferred": False}, + {"name": "BrowserWait", "description": "Wait for page loads or animations", "category": "browser_action", "deferred": False}, +] diff --git a/backend/apps/tools/discover_tools/DiscoveryError.py b/backend/apps/tools/discover_tools/DiscoveryError.py new file mode 100644 index 00000000..09cc66f0 --- /dev/null +++ b/backend/apps/tools/discover_tools/DiscoveryError.py @@ -0,0 +1,8 @@ +class DiscoveryError(Exception): + """Raised when MCP tool discovery fails.""" + pass + + +class DiscoveryConfigError(DiscoveryError): + """Raised for invalid discovery configuration (client error).""" + pass \ No newline at end of file diff --git a/backend/apps/tools/discover_tools/discover_tools.py b/backend/apps/tools/discover_tools/discover_tools.py new file mode 100644 index 00000000..5758be02 --- /dev/null +++ b/backend/apps/tools/discover_tools/discover_tools.py @@ -0,0 +1,42 @@ +from typing import Any +from backend.apps.tools.discover_tools.DiscoveryError import DiscoveryError, DiscoveryConfigError +from backend.apps.tools.discover_tools.utils.discover_mcp_tools_stdio import discover_mcp_tools_stdio +from backend.apps.tools.discover_tools.utils.discover_mcp_tools_http import discover_mcp_tools_http +from backend.apps.tools.discover_tools.utils.discover_mcp_tools_sse import discover_mcp_tools_sse +from typeguard import typechecked + +# TODO: better type specing of this whole func +@typechecked +async def discover_tools(config: dict[str, Any], tool_name: str = "") -> list[dict]: + """Probe an MCP server using the appropriate transport and return discovered tools. + + config is the raw mcp_config dict from a ToolDefinition (with credentials + already injected by the converter if needed). + + Returns a list of dicts with keys: name, description, inputSchema. + """ + transport = config.get("type", "") + + if transport == "stdio": + command = config.get("command", "") + if not command: + raise DiscoveryConfigError("stdio transport requires a 'command' in MCP config") + return await discover_mcp_tools_stdio( + command=command, + args=config.get("args"), + env=config.get("env"), + ) + + if transport in ("http", "sse") or config.get("url"): + url = config.get("url", "") + if not url: + raise DiscoveryConfigError("HTTP/SSE transport requires a 'url' in MCP config") + if transport == "sse": + return await discover_mcp_tools_sse(url, config.get("headers")) + try: + return await discover_mcp_tools_http(url, config.get("headers")) + except DiscoveryError: + print(f"[discover_tools] Streamable HTTP failed for {tool_name}, retrying with SSE") + return await discover_mcp_tools_sse(url, config.get("headers")) + + raise DiscoveryConfigError(f"Unsupported MCP transport type: '{transport}'") diff --git a/backend/apps/tools/discover_tools/utils/discover_mcp_tools_http.py b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_http.py new file mode 100644 index 00000000..baea3140 --- /dev/null +++ b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_http.py @@ -0,0 +1,70 @@ +import json +from typing import Optional +import httpx +from typeguard import typechecked +from backend.apps.tools.discover_tools.DiscoveryError import DiscoveryError + +# TODO: better type specing of return value +@typechecked +def p_parse_sse_json(text: str) -> Optional[dict]: + for line in text.splitlines(): + stripped: str = line.strip() + if stripped.startswith("data:"): + payload: str = stripped[len("data:"):].strip() + if payload: + try: + return json.loads(payload) + except json.JSONDecodeError: + continue + try: + return json.loads(text) + except json.JSONDecodeError: + return None + + +# TODO: better type specing throughout this whole func +@typechecked +async def discover_mcp_tools_http(url: str, headers: dict | None = None) -> list[dict]: + """Discover tools via streamable HTTP JSON-RPC.""" + h: dict[str, str] = { + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + **(headers or {}), + } + async with httpx.AsyncClient(timeout=30.0) as client: + init_resp: httpx.Response = await client.post(url, headers=h, json={ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": { + "protocolVersion": "2025-03-26", + "capabilities": {}, + "clientInfo": {"name": "openswarm", "version": "0.1.0"}, + }, + }) + if init_resp.status_code not in (200, 201): + raise DiscoveryError(f"MCP initialize failed: {init_resp.status_code}") + + session_id = init_resp.headers.get("mcp-session-id", "") + if session_id: + h["mcp-session-id"] = session_id + + await client.post(url, headers=h, json={ + "jsonrpc": "2.0", "method": "notifications/initialized", + }) + + list_resp = await client.post(url, headers=h, json={ + "jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}, + }) + if list_resp.status_code not in (200, 201): + raise DiscoveryError(f"MCP tools/list failed: {list_resp.status_code}") + + ct = list_resp.headers.get("content-type", "") + data: Optional[dict] = p_parse_sse_json(list_resp.text) if "text/event-stream" in ct else list_resp.json() + + if not data: + raise DiscoveryError("Empty response from MCP server") + + tools_list = data.get("result", {}).get("tools", []) + return [ + {"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} + for t in tools_list + ] \ No newline at end of file diff --git a/backend/apps/tools/discover_tools/utils/discover_mcp_tools_sse.py b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_sse.py new file mode 100644 index 00000000..f6c8ea2a --- /dev/null +++ b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_sse.py @@ -0,0 +1,26 @@ +from mcp.client.sse import sse_client +from mcp import ClientSession +from mcp.types import Implementation +from exceptiongroup import BaseExceptionGroup + +from backend.apps.tools.discover_tools.DiscoveryError import DiscoveryError +from typeguard import typechecked + +@typechecked +async def discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list[dict]: + """Discover tools via SSE transport using the mcp SDK client.""" + try: + async with sse_client(url=url, headers=headers, timeout=30, sse_read_timeout=30) as (read_stream, write_stream): + async with ClientSession( + read_stream, write_stream, + client_info=Implementation(name="openswarm", version="0.1.0"), + ) as session: + await session.initialize() + result = await session.list_tools() + return [ + {"name": t.name, "description": t.description or "", "inputSchema": t.inputSchema if t.inputSchema else None} + for t in result.tools + ] + except BaseExceptionGroup as eg: + first = eg.exceptions[0] if eg.exceptions else eg + raise DiscoveryError(f"SSE discovery failed: {first}") from first \ No newline at end of file diff --git a/backend/apps/tools/discover_tools/utils/discover_mcp_tools_stdio.py b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_stdio.py new file mode 100644 index 00000000..207398d7 --- /dev/null +++ b/backend/apps/tools/discover_tools/utils/discover_mcp_tools_stdio.py @@ -0,0 +1,97 @@ +import asyncio +import json +import os + +from backend.apps.tools.shared_utils.mcp_config import resolve_command, augmented_path +from backend.apps.tools.discover_tools.DiscoveryError import DiscoveryError, DiscoveryConfigError +from typeguard import typechecked + +# TODO: better type specing of this whole func +@typechecked +async def discover_mcp_tools_stdio( + command: str, + args: list[str] | None = None, + env: dict | None = None, +) -> list[dict]: + """Discover tools by spawning an MCP server subprocess via stdio.""" + cmd_path = resolve_command(command) + if not cmd_path: + raise DiscoveryConfigError(f"Command '{command}' not found on PATH or common install locations") + + proc_env = {**os.environ, **(env or {}), "PATH": augmented_path()} + proc_env.pop("PYTHONPATH", None) + + proc = await asyncio.create_subprocess_exec( + cmd_path, *(args or []), + stdin=asyncio.subprocess.PIPE, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=proc_env, + limit=1024 * 1024, + ) + + async def _send(msg: dict) -> None: + assert proc.stdin is not None + line = json.dumps(msg) + "\n" + proc.stdin.write(line.encode()) + await proc.stdin.drain() + + async def _recv() -> dict: + assert proc.stdout is not None and proc.stderr is not None + while True: + line = await asyncio.wait_for(proc.stdout.readline(), timeout=30.0) + if not line: + stderr_out = "" + try: + stderr_out = (await asyncio.wait_for(proc.stderr.read(4096), timeout=2.0)).decode(errors="replace") + except (asyncio.TimeoutError, Exception): + pass + raise DiscoveryError( + f"MCP stdio process exited unexpectedly{': ' + stderr_out if stderr_out else ''}" + ) + stripped = line.decode(errors="replace").strip() + if not stripped: + continue + try: + data = json.loads(stripped) + except json.JSONDecodeError: + continue + if "id" in data: + return data + + try: + await _send({ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": { + "protocolVersion": "2025-03-26", + "capabilities": {}, + "clientInfo": {"name": "openswarm", "version": "0.1.0"}, + }, + }) + await _recv() + await _send({"jsonrpc": "2.0", "method": "notifications/initialized"}) + await _send({"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}}) + data = await _recv() + tools_list = data.get("result", {}).get("tools", []) + return [ + {"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} + for t in tools_list + ] + except (DiscoveryError, DiscoveryConfigError): + raise + except asyncio.TimeoutError: + raise DiscoveryError("MCP stdio server timed out during discovery") + finally: + try: + if proc.stdin: + proc.stdin.close() + except Exception: + pass + try: + proc.terminate() + await asyncio.wait_for(proc.wait(), timeout=5.0) + except Exception: + try: + proc.kill() + except Exception: + pass diff --git a/backend/apps/tools/oauth/oauth.py b/backend/apps/tools/oauth/oauth.py new file mode 100644 index 00000000..34290115 --- /dev/null +++ b/backend/apps/tools/oauth/oauth.py @@ -0,0 +1,265 @@ +"""OAuth flow logic — callback, start, disconnect, refresh. + +The tool store is injected via set_store() from the tools sub-app +to avoid circular imports. +""" + +import base64 +import hashlib +import logging +import os +import secrets +import time +from typing import Any, Optional +from urllib.parse import urlencode + +import httpx +from fastapi import HTTPException, Query +from fastapi.responses import HTMLResponse + +from backend.apps.tools.oauth.oauth_providers import resolve_oauth_provider +from backend.core.db.PydanticStore import PydanticStore +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from backend.ports import BACKEND_DEV_PORT + +logger = logging.getLogger(__name__) + +_pending_oauth: dict[str, str] = {} +_pending_pkce: dict[str, str] = {} + +_store: Optional[PydanticStore[ToolDefinition]] = None + + +def set_store(store: PydanticStore[ToolDefinition]) -> None: + global _store + _store = store + + +def _get_store() -> PydanticStore[ToolDefinition]: + assert _store is not None, "OAuth store not initialized — call set_store() first" + return _store + + +async def oauth_callback(code: str = Query(...), state: str = Query("")) -> HTMLResponse: + tool_id = _pending_oauth.pop(state, None) + if not tool_id: + tool_id = _pending_oauth.pop(state.split(":")[-1] if ":" in state else state, None) + if not tool_id: + return HTMLResponse("

Invalid OAuth state

", status_code=400) + + store = _get_store() + tool = store.load(tool_id) + provider = resolve_oauth_provider(tool.oauth_provider) + + client_id = os.environ.get(provider.client_id_env, "") + client_secret = os.environ.get(provider.client_secret_env, "") + port = os.environ.get("OPENSWARM_PORT", str(BACKEND_DEV_PORT)) + redirect_uri = f"http://localhost:{port}/api/tools/oauth/callback" + + token_data: dict[str, str] = { + "code": code, "redirect_uri": redirect_uri, "grant_type": "authorization_code", + } + headers: dict[str, str] = {} + + if provider.token_auth_method == "basic": + creds = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode() + headers["Authorization"] = f"Basic {creds}" + elif provider.token_auth_method == "basic_json": + creds = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode() + headers["Authorization"] = f"Basic {creds}" + headers["Content-Type"] = "application/json" + else: + token_data["client_id"] = client_id + token_data["client_secret"] = client_secret + + if (tool.oauth_provider or "google") == "github": + headers["Accept"] = "application/json" + + code_verifier = _pending_pkce.pop(state, None) + if code_verifier: + token_data["code_verifier"] = code_verifier + + async with httpx.AsyncClient(timeout=15.0) as client: + if provider.token_auth_method == "basic_json": + resp = await client.post(provider.token_url, json=token_data, headers=headers) + else: + resp = await client.post(provider.token_url, data=token_data, headers=headers) + + if resp.status_code != 200: + logger.warning(f"OAuth token exchange failed: {resp.text}") + return HTMLResponse( + f"

Token exchange failed

{resp.text}
", + status_code=400, + ) + + tokens = resp.json() + + access_token = tokens.get("access_token", "") + if provider.token_response_path and not access_token: + obj: Any = tokens + for part in provider.token_response_path.split("."): + obj = obj.get(part, {}) if isinstance(obj, dict) else "" + if isinstance(obj, str) and obj: + access_token = obj + + tool.oauth_tokens = { + "access_token": access_token, + "refresh_token": tokens.get("refresh_token", ""), + "token_expiry": time.time() + tokens.get("expires_in", 3600), + } + + for response_path, env_var in provider.extra_token_fields.items(): + obj_val: Any = tokens + for part in response_path.split("."): + obj_val = obj_val.get(part, "") if isinstance(obj_val, dict) else "" + if obj_val: + tool.oauth_tokens[env_var] = str(obj_val) + + tool.auth_status = "connected" + + if access_token and provider.userinfo_url: + try: + async with httpx.AsyncClient(timeout=10.0) as info_client: + info_resp = await info_client.get( + provider.userinfo_url, + headers={"Authorization": f"Bearer {access_token}"}, + ) + if info_resp.status_code == 200: + tool.connected_account_email = info_resp.json().get(provider.userinfo_field) + except Exception as e: + logger.warning(f"Failed to fetch userinfo for {tool.oauth_provider or 'google'}: {e}") + + if (tool.oauth_provider or "google") == "notion" and not tool.connected_account_email: + workspace_name = tokens.get("workspace_name") + if workspace_name: + tool.connected_account_email = workspace_name + + store.save(tool) + + return HTMLResponse( + "" + '

Connected successfully!

' + '

You can close this window.

' + "" + "" + ) + + +async def oauth_start(tool_id: str) -> dict: + store = _get_store() + tool = store.load(tool_id) + provider = resolve_oauth_provider(tool.oauth_provider) + + client_id = os.environ.get(provider.client_id_env, "") + if not client_id: + raise HTTPException(status_code=400, detail=f"{provider.client_id_env} not set in backend .env") + + port = os.environ.get("OPENSWARM_PORT", str(BACKEND_DEV_PORT)) + redirect_uri = f"http://localhost:{port}/api/tools/oauth/callback" + provider_key = tool.oauth_provider or "google" + state = f"{provider_key}:{tool_id}" + + _pending_oauth[state] = tool_id + + params: dict[str, str] = { + "client_id": client_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "state": state, + **provider.extra_auth_params, + } + if provider.scopes: + params["scope"] = " ".join(provider.scopes) + + if provider.pkce_required: + code_verifier = secrets.token_urlsafe(64) + code_challenge = base64.urlsafe_b64encode( + hashlib.sha256(code_verifier.encode()).digest() + ).rstrip(b"=").decode() + params["code_challenge"] = code_challenge + params["code_challenge_method"] = "S256" + _pending_pkce[state] = code_verifier + + auth_url = f"{provider.auth_url}?{urlencode(params)}" + return {"auth_url": auth_url} + + +async def oauth_disconnect(tool_id: str) -> dict: + store = _get_store() + tool = store.load(tool_id) + access_token = tool.oauth_tokens.get("access_token") + + if access_token: + provider = resolve_oauth_provider(tool.oauth_provider) + revoke_url = provider.revoke_url or "https://oauth2.googleapis.com/revoke" + try: + async with httpx.AsyncClient(timeout=10.0) as client: + await client.post( + revoke_url, + params={"token": access_token}, + headers={"Content-Type": "application/x-www-form-urlencoded"}, + ) + except Exception as e: + logger.warning(f"Failed to revoke token for tool {tool.id}: {e}") + + tool.oauth_tokens = {} + tool.auth_status = "configured" + tool.connected_account_email = None + store.save(tool) + return {"ok": True, "tool": tool.model_dump()} + + +async def refresh_oauth_token(tool: ToolDefinition) -> Optional[str]: + """Refresh an expired OAuth token. Returns the fresh access_token or None. + + Mutates the tool in-place and saves to the store if refresh succeeds. + """ + if tool.auth_type != "oauth2": + return None + refresh_token = tool.oauth_tokens.get("refresh_token") + if not refresh_token: + return None + expiry = tool.oauth_tokens.get("token_expiry", 0) + if time.time() < expiry - 60: + return tool.oauth_tokens.get("access_token") + + provider = resolve_oauth_provider(tool.oauth_provider) + client_id = os.environ.get(provider.client_id_env, "") + client_secret = os.environ.get(provider.client_secret_env, "") + if not client_id or not client_secret: + return None + + try: + async with httpx.AsyncClient(timeout=15.0) as client: + resp = await client.post(provider.token_url, data={ + "client_id": client_id, + "client_secret": client_secret, + "refresh_token": refresh_token, + "grant_type": "refresh_token", + }) + if resp.status_code == 200: + data = resp.json() + new_token = data["access_token"] + tool.oauth_tokens["access_token"] = new_token + tool.oauth_tokens["token_expiry"] = time.time() + data.get("expires_in", 3600) + + if not tool.connected_account_email and provider.userinfo_url: + try: + async with httpx.AsyncClient(timeout=10.0) as info_client: + info_resp = await info_client.get( + provider.userinfo_url, + headers={"Authorization": f"Bearer {new_token}"}, + ) + if info_resp.status_code == 200: + tool.connected_account_email = info_resp.json().get(provider.userinfo_field) + except Exception: + pass + + _get_store().save(tool) + return new_token + except Exception as e: + logger.warning(f"OAuth token refresh failed for tool {tool.id}: {e}") + return None diff --git a/backend/apps/tools/oauth/oauth_providers.py b/backend/apps/tools/oauth/oauth_providers.py new file mode 100644 index 00000000..0aa5fde6 --- /dev/null +++ b/backend/apps/tools/oauth/oauth_providers.py @@ -0,0 +1,181 @@ +"""OAuth provider definitions — pure data, no route handlers.""" + +import os +from dataclasses import dataclass, field + + +# TODO: wtf is this, bruh we gotta remove this shit asap r u fr rn. Unacceptable. +os.environ.setdefault("GOOGLE_OAUTH_CLIENT_ID", "6741219524-8vpt07arcc5rvkdb4j1b6v9g53469ugq.apps.googleusercontent.com") +os.environ.setdefault("GOOGLE_OAUTH_CLIENT_SECRET", "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") +os.environ.setdefault("GITHUB_OAUTH_CLIENT_ID", "Ov23liDcwNJaKMjXY2jI") +os.environ.setdefault("GITHUB_OAUTH_CLIENT_SECRET", "b25fe39409896aad3fd5155f032e9868440002f8") +os.environ.setdefault("SLACK_CLIENT_ID", "10795695056323.10799999254534") +os.environ.setdefault("SLACK_CLIENT_SECRET", "d3a85a286bb0205157d7e4963502a91d") +os.environ.setdefault("FIGMA_CLIENT_ID", "q6WduT7UuPaO6lM88v6ddN") +os.environ.setdefault("FIGMA_CLIENT_SECRET", "dhNZdbEuyEWC15cKLwWpqTclyOSplD") +os.environ.setdefault("AIRTABLE_CLIENT_ID", "0699038b-a3a4-46b2-8fa6-690eb76fadfa") +os.environ.setdefault("AIRTABLE_CLIENT_SECRET", "187fa83c8bab8ebcd11b8f226d75e7a1f14a8174ac0494463c1a53e66a3036d0") +os.environ.setdefault("HUBSPOT_CLIENT_ID", "6f4a1d4c-6a2f-4336-9b65-2cd84e218ff6") +os.environ.setdefault("HUBSPOT_CLIENT_SECRET", "5747b5de-0800-4c35-a2da-e0655ee7ea37") + + +@dataclass +class OAuthProvider: + auth_url: str + token_url: str + scopes: list[str] + userinfo_url: str | None + userinfo_field: str + client_id_env: str + client_secret_env: str + token_env_mapping: dict[str, str] + extra_auth_params: dict[str, str] = field(default_factory=dict) + revoke_url: str | None = None + token_response_path: str | None = None + token_auth_method: str = "form" + pkce_required: bool = False + env_value_transform: str | None = None + extra_token_fields: dict[str, str] = field(default_factory=dict) + + +OAUTH_PROVIDERS: dict[str, OAuthProvider] = { + "google": OAuthProvider( + auth_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=[ + "openid", + "https://www.googleapis.com/auth/userinfo.email", + "https://www.googleapis.com/auth/gmail.modify", + "https://www.googleapis.com/auth/calendar", + "https://www.googleapis.com/auth/drive", + "https://www.googleapis.com/auth/contacts.readonly", + ], + userinfo_url="https://www.googleapis.com/oauth2/v2/userinfo", + userinfo_field="email", + client_id_env="GOOGLE_OAUTH_CLIENT_ID", + client_secret_env="GOOGLE_OAUTH_CLIENT_SECRET", + token_env_mapping={ + "access_token": "OAUTH_ACCESS_TOKEN", + "refresh_token": "GOOGLE_WORKSPACE_REFRESH_TOKEN", + "_client_id": "GOOGLE_WORKSPACE_CLIENT_ID", + "_client_secret": "GOOGLE_WORKSPACE_CLIENT_SECRET", + }, + extra_auth_params={"access_type": "offline", "prompt": "consent"}, + ), + "github": OAuthProvider( + auth_url="https://github.com/login/oauth/authorize", + token_url="https://github.com/login/oauth/access_token", + scopes=["repo", "read:user", "user:email"], + userinfo_url="https://api.github.com/user", + userinfo_field="login", + client_id_env="GITHUB_OAUTH_CLIENT_ID", + client_secret_env="GITHUB_OAUTH_CLIENT_SECRET", + token_env_mapping={"access_token": "GITHUB_PERSONAL_ACCESS_TOKEN"}, + ), + "slack": OAuthProvider( + auth_url="https://slack.com/oauth/v2/authorize", + token_url="https://slack.com/api/oauth.v2.access", + scopes=[ + "channels:read", "channels:history", "chat:write", + "groups:read", "groups:history", "im:read", "im:history", + "mpim:read", "mpim:history", "users:read", "users:read.email", + "team:read", "reactions:read", "reactions:write", + "files:read", "files:write", + ], + userinfo_url="https://slack.com/api/auth.test", + userinfo_field="user", + client_id_env="SLACK_CLIENT_ID", + client_secret_env="SLACK_CLIENT_SECRET", + token_env_mapping={"access_token": "SLACK_BOT_TOKEN"}, + extra_token_fields={"team.id": "SLACK_TEAM_ID"}, + ), + "notion": OAuthProvider( + auth_url="https://api.notion.com/v1/oauth/authorize", + token_url="https://api.notion.com/v1/oauth/token", + scopes=[], + userinfo_url=None, + userinfo_field="owner", + client_id_env="NOTION_OAUTH_CLIENT_ID", + client_secret_env="NOTION_OAUTH_CLIENT_SECRET", + token_env_mapping={"access_token": "OPENAPI_MCP_HEADERS"}, + extra_auth_params={"owner": "user"}, + token_auth_method="basic_json", + env_value_transform="notion_headers", + ), + "spotify": OAuthProvider( + auth_url="https://accounts.spotify.com/authorize", + token_url="https://accounts.spotify.com/api/token", + scopes=[ + "user-read-playback-state", "user-modify-playback-state", + "user-read-currently-playing", "playlist-read-private", + "playlist-modify-public", "playlist-modify-private", + "user-library-read", "user-library-modify", + "user-read-recently-played", "user-top-read", + ], + userinfo_url="https://api.spotify.com/v1/me", + userinfo_field="display_name", + client_id_env="SPOTIFY_CLIENT_ID", + client_secret_env="SPOTIFY_CLIENT_SECRET", + token_env_mapping={ + "access_token": "SPOTIFY_ACCESS_TOKEN", + "refresh_token": "SPOTIFY_REFRESH_TOKEN", + "_client_id": "SPOTIFY_CLIENT_ID", + "_client_secret": "SPOTIFY_CLIENT_SECRET", + }, + token_auth_method="basic", + ), + "figma": OAuthProvider( + auth_url="https://www.figma.com/oauth", + token_url="https://api.figma.com/v1/oauth/token", + scopes=[ + "current_user:read", "file_content:read", "file_metadata:read", + "file_comments:read", "file_comments:write", + "file_versions:read", "file_variables:read", + ], + userinfo_url="https://api.figma.com/v1/me", + userinfo_field="email", + client_id_env="FIGMA_CLIENT_ID", + client_secret_env="FIGMA_CLIENT_SECRET", + token_env_mapping={"access_token": "FIGMA_API_KEY"}, + ), + "airtable": OAuthProvider( + auth_url="https://airtable.com/oauth2/v1/authorize", + token_url="https://airtable.com/oauth2/v1/token", + scopes=[ + "data.records:read", "data.records:write", + "data.recordComments:read", "data.recordComments:write", + "schema.bases:read", "schema.bases:write", + "user.email:read", "webhook:manage", + ], + userinfo_url="https://api.airtable.com/v0/meta/whoami", + userinfo_field="email", + client_id_env="AIRTABLE_CLIENT_ID", + client_secret_env="AIRTABLE_CLIENT_SECRET", + token_env_mapping={"access_token": "AIRTABLE_API_KEY"}, + pkce_required=True, + token_auth_method="basic", + ), + "hubspot": OAuthProvider( + auth_url="https://mcp-na2.hubspot.com/oauth/authorize/user", + token_url="https://api.hubapi.com/oauth/v1/token", + scopes=[], + userinfo_url=None, + userinfo_field="user", + client_id_env="HUBSPOT_CLIENT_ID", + client_secret_env="HUBSPOT_CLIENT_SECRET", + token_env_mapping={ + "access_token": "PRIVATE_APP_ACCESS_TOKEN", + "refresh_token": "HUBSPOT_REFRESH_TOKEN", + }, + pkce_required=True, + ), +} + + +def resolve_oauth_provider(oauth_provider_key: str | None) -> OAuthProvider: + """Resolve the OAuth provider by key, defaulting to Google.""" + key = oauth_provider_key or "google" + provider = OAUTH_PROVIDERS.get(key) + if not provider: + raise ValueError(f"Unknown OAuth provider: {key}") + return provider diff --git a/backend/apps/tools/shared_utils/ToolDefinition.py b/backend/apps/tools/shared_utils/ToolDefinition.py new file mode 100644 index 00000000..83268bc3 --- /dev/null +++ b/backend/apps/tools/shared_utils/ToolDefinition.py @@ -0,0 +1,23 @@ +from pydantic import BaseModel, Field +from typing import Optional, Any +from uuid import uuid4 +from backend.core.tools.shared_structs.TOOL_PERMISSIONS import TOOL_PERMISSIONS + + +class ToolDefinition(BaseModel): + model_config = {"extra": "ignore"} + + id: str = Field(default_factory=lambda: uuid4().hex) + name: str + description: str = "" + command: str = "" + mcp_config: dict[str, Any] = Field(default_factory=dict) + credentials: dict[str, str] = Field(default_factory=dict) + auth_type: str = "none" + auth_status: str = "none" + oauth_provider: Optional[str] = None + oauth_tokens: dict[str, Any] = Field(default_factory=dict) + tool_permissions: dict[str, TOOL_PERMISSIONS] = Field(default_factory=dict) + tool_descriptions: dict[str, str] = Field(default_factory=dict) # tool_descriptions[tool_name] = tool_description + connected_account_email: Optional[str] = None + enabled: bool = True diff --git a/backend/apps/tools/shared_utils/mcp_config.py b/backend/apps/tools/shared_utils/mcp_config.py new file mode 100644 index 00000000..5b50105d --- /dev/null +++ b/backend/apps/tools/shared_utils/mcp_config.py @@ -0,0 +1,69 @@ +"""PATH resolution and command lookup helpers for MCP stdio servers. + +These are needed because packaged Electron apps and various Node version +managers (nvm, fnm, volta) install binaries in non-standard locations that +may not be on the default PATH. +""" + +import os +import shutil +from typing import List, Optional + +# TODO: either remove the import, or make it a non private var +from backend.config.paths import P_BACKEND_DIR +from typeguard import typechecked + +P_UV_BIN_DIR = os.path.join(P_BACKEND_DIR, "uv-bin") + +@typechecked +def p_extra_bin_dirs() -> list[str]: + """Well-known user-local bin directories that may not be on PATH in packaged apps.""" + home = os.path.expanduser("~") + dirs: List[str] = [ + P_UV_BIN_DIR, + os.path.join(home, ".bun", "bin"), + os.path.join(home, ".cargo", "bin"), + os.path.join(home, ".local", "bin"), + os.path.join(home, ".volta", "bin"), + "/opt/homebrew/bin", + "/usr/local/bin", + ] + nvm_node = os.path.join(home, ".nvm", "versions", "node") + try: + if os.path.isdir(nvm_node): + versions = sorted(os.listdir(nvm_node), reverse=True) + if versions: + dirs.insert(0, os.path.join(nvm_node, versions[0], "bin")) + except OSError: + pass + fnm_bin = os.path.join(home, "Library", "Application Support", "fnm", "aliases", "default", "bin") + if os.path.isdir(fnm_bin): + dirs.insert(0, fnm_bin) + return dirs + + +@typechecked +def resolve_command(command: str) -> Optional[str]: + """Find an executable on PATH or well-known bin directories.""" + found = shutil.which(command) + if found: + return found + for d in p_extra_bin_dirs(): + candidate = os.path.join(d, command) + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + return candidate + return None + + +@typechecked +def augmented_path() -> str: + """Build a PATH string that includes well-known extra bin dirs.""" + extra = [d for d in p_extra_bin_dirs() if os.path.isdir(d)] + current = os.environ.get("PATH", "") + seen: set[str] = set[str]() + parts: list[str] = [] + for p in extra + current.split(os.pathsep): + if p and p not in seen: + seen.add(p) + parts.append(p) + return os.pathsep.join(parts) diff --git a/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_http_sse_tool.py b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_http_sse_tool.py new file mode 100644 index 00000000..54201219 --- /dev/null +++ b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_http_sse_tool.py @@ -0,0 +1,22 @@ +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from backend.core.tools.shared_structs.MCP_Tool import SSE_HTTP_MCP_Tool +from typeguard import typechecked + +# TODO: better type specing of this whole func +@typechecked +def build_http_sse_tool( + tool_def: ToolDefinition, + config: dict, + server_name: str, + transport: str, +) -> SSE_HTTP_MCP_Tool: + return SSE_HTTP_MCP_Tool( + name=tool_def.name, + description=tool_def.description, + deferred=False, + permission="ask", + server_name=server_name, + transport=transport, # type: ignore[arg-type] + url=config.get("url"), + headers=config.get("headers", {}), + ) diff --git a/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_stdio_tool.py b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_stdio_tool.py new file mode 100644 index 00000000..888c1b76 --- /dev/null +++ b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/build_stdio_tool.py @@ -0,0 +1,49 @@ +import os + +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from backend.core.tools.shared_structs.MCP_Tool import STDIO_MCP_Tool +from backend.apps.tools.shared_utils.mcp_config import resolve_command, augmented_path +# TODO: either remove the import, or make them non private vars +from backend.config.paths import P_BACKEND_DIR, p_is_packaged +from typeguard import typechecked + + +# TODO: better type specing of this whole func +@typechecked +def build_stdio_tool( + tool_def: ToolDefinition, + config: dict, + server_name: str, +) -> STDIO_MCP_Tool: + command = config.get("command", "") + if command: + resolved = resolve_command(command) + if resolved: + command = resolved + else: + print(f"[build_stdio_tool] Command '{command}' not found on PATH or bundled directories") + + env = config.get("env", {}) + env.setdefault("PATH", augmented_path()) + env.setdefault("PYTHONPATH", "") + + if p_is_packaged: + resources = os.path.dirname(os.path.dirname(P_BACKEND_DIR)) + bundled_python = os.path.join(resources, "python-env", "bin", "python3") + if os.path.exists(bundled_python): + env.setdefault("UV_PYTHON", bundled_python) + else: + venv_python = os.path.join(P_BACKEND_DIR, ".venv", "bin", "python3") + if os.path.exists(venv_python): + env.setdefault("UV_PYTHON", venv_python) + + return STDIO_MCP_Tool( + name=tool_def.name, + description=tool_def.description, + deferred=False, + permission="ask", + server_name=server_name, + command=command or None, + args=config.get("args", []), + env=env, + ) \ No newline at end of file diff --git a/backend/apps/tools/tool_definition_to_mcp_tool/helpers/inject_credentials.py b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/inject_credentials.py new file mode 100644 index 00000000..ad3166bf --- /dev/null +++ b/backend/apps/tools/tool_definition_to_mcp_tool/helpers/inject_credentials.py @@ -0,0 +1,67 @@ +import json +import os +from typing import Optional + +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from typeguard import typechecked + +# TODO: better type specing of this whole func +@typechecked +def inject_credentials( + tool_def: ToolDefinition, + config: dict, + oauth_providers: Optional[dict], +) -> None: + """Mutate config in-place to inject credentials and OAuth tokens.""" + transport = config.get("type", "") + + if tool_def.credentials: + if transport in ("http", "sse"): + headers = config.setdefault("headers", {}) + for key, val in tool_def.credentials.items(): + if key.lower() in ("authorization", "api_key", "api-key"): + headers.setdefault("Authorization", f"Bearer {val}") + else: + env = config.setdefault("env", {}) + env.update(tool_def.credentials) + + if tool_def.auth_type != "oauth2" or not tool_def.oauth_tokens.get("access_token"): + return + + if transport in ("http", "sse"): + headers = config.setdefault("headers", {}) + headers["Authorization"] = f"Bearer {tool_def.oauth_tokens['access_token']}" + return + + env = config.setdefault("env", {}) + provider_key = tool_def.oauth_provider or "google" + provider = (oauth_providers or {}).get(provider_key) + + if not provider: + env["OAUTH_ACCESS_TOKEN"] = tool_def.oauth_tokens["access_token"] + return + + for token_field, env_var in provider.token_env_mapping.items(): + if token_field.startswith("_client_id"): + val = os.environ.get(provider.client_id_env, "") + elif token_field.startswith("_client_secret"): + val = os.environ.get(provider.client_secret_env, "") + else: + val = tool_def.oauth_tokens.get(token_field, "") + if val: + if provider.env_value_transform == "notion_headers" and token_field == "access_token": + val = json.dumps({ + "Authorization": f"Bearer {val}", + "Notion-Version": "2022-06-28", + }) + env[env_var] = val + + for _, env_var in provider.extra_token_fields.items(): + val = tool_def.oauth_tokens.get(env_var, "") + if val: + env[env_var] = val + + if provider_key == "figma" and tool_def.oauth_tokens.get("access_token"): + args = config.get("args", []) + if "--figma-api-key" not in args: + config["args"] = args + ["--figma-api-key", tool_def.oauth_tokens["access_token"]] \ No newline at end of file diff --git a/backend/apps/tools/tool_definition_to_mcp_tool/tool_definition_to_mcp_tool.py b/backend/apps/tools/tool_definition_to_mcp_tool/tool_definition_to_mcp_tool.py new file mode 100644 index 00000000..0776e809 --- /dev/null +++ b/backend/apps/tools/tool_definition_to_mcp_tool/tool_definition_to_mcp_tool.py @@ -0,0 +1,52 @@ +"""Convert a persisted ToolDefinition into a typed MCP_Tool for the Toolkit tree. + +Replaces the legacy derive_mcp_config() — instead of producing raw dicts, +we produce STDIO_MCP_Tool or SSE_HTTP_MCP_Tool instances that slot directly +into the Toolkit and participate in collect_mcp_servers / collect_tool_permissions. +""" + +from typing import Optional + +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from backend.core.tools.shared_structs.MCP_Tool import MCP_Tool +from backend.apps.tools.tool_definition_to_mcp_tool.helpers.inject_credentials import inject_credentials +from backend.apps.tools.tool_definition_to_mcp_tool.helpers.build_stdio_tool import build_stdio_tool +from backend.apps.tools.tool_definition_to_mcp_tool.helpers.build_http_sse_tool import build_http_sse_tool +from typeguard import typechecked +import re + + +P_SANITIZE_RE: re.Pattern[str] = re.compile(r"[^a-zA-Z0-9\-]") + +@typechecked +def p_sanitize_mcp_server_name(name: str) -> str: + return P_SANITIZE_RE.sub("-", name).strip("-").lower() + + +# TODO: better type specing of this whole func +@typechecked +def tool_definition_to_mcp_tool( + tool_def: ToolDefinition, + oauth_providers: Optional[dict] = None, +) -> Optional[MCP_Tool]: + """Build an MCP_Tool from a ToolDefinition, or None if config is missing. + + oauth_providers is the OAUTH_PROVIDERS dict from the oauth module, + passed in to avoid a circular import. + """ + if not tool_def.mcp_config: + return None + + config: dict = dict(tool_def.mcp_config) + transport = config.get("type", "") + server_name = p_sanitize_mcp_server_name(tool_def.name) + + inject_credentials(tool_def, config, oauth_providers) + + if transport == "stdio": + return build_stdio_tool(tool_def, config, server_name) + elif transport in ("http", "sse"): + return build_http_sse_tool(tool_def, config, server_name, transport) + + print(f"[tool_definition_to_mcp_tool] Unsupported MCP transport type '{transport}' for tool {tool_def.name}") + return None diff --git a/backend/apps/tools/tools.py b/backend/apps/tools/tools.py new file mode 100644 index 00000000..507acbce --- /dev/null +++ b/backend/apps/tools/tools.py @@ -0,0 +1,255 @@ +"""Tools sub-app — CRUD for user-installed MCP tools, builtin permissions, and discovery.""" + +from claude_agent_sdk.types import McpServerConfig +from pydantic import BaseModel, Field +import json +import logging +import os +import time +from contextlib import asynccontextmanager +from typing import Any, Optional + +from fastapi import HTTPException + +from backend.config.Apps import SubApp +from backend.config.paths import DB_ROOT +from backend.core.db.PydanticStore import PydanticStore +from backend.apps.tools.shared_utils.ToolDefinition import ToolDefinition +from backend.apps.tools.discover_tools.discover_tools import discover_tools +from backend.apps.tools.discover_tools.DiscoveryError import DiscoveryError, DiscoveryConfigError +from backend.apps.tools.tool_definition_to_mcp_tool.tool_definition_to_mcp_tool import tool_definition_to_mcp_tool +from backend.apps.tools.oauth import oauth +from backend.apps.tools.oauth.oauth import refresh_oauth_token +from backend.apps.tools.oauth.oauth_providers import OAUTH_PROVIDERS +from backend.apps.tools.builtin_tools import BUILTIN_TOOLS +from backend.core.tools.shared_structs.TOOL_PERMISSIONS import TOOL_PERMISSIONS +from backend.core.tools.shared_structs.Toolkit import Toolkit +from backend.core.tools.shared_structs.Tool import Tool +from backend.core.tools.shared_structs.MCP_Tool import MCP_Tool + +logger = logging.getLogger(__name__) + +TOOLS_DIR = os.path.join(DB_ROOT, "tools") +BUILTIN_PERMS_PATH = os.path.join(TOOLS_DIR, "builtin_permissions.json") + +TOOL_STORE: PydanticStore[ToolDefinition] = PydanticStore[ToolDefinition]( + model_cls=ToolDefinition, + data_dir=TOOLS_DIR, + id_field="id", + dump_mode="json", + not_found_detail="Tool not found", +) + + +@asynccontextmanager +async def tools_lifespan(): + os.makedirs(TOOLS_DIR, exist_ok=True) + oauth.set_store(TOOL_STORE) + yield + + +tools = SubApp("tools", tools_lifespan) + + +# --------------------------------------------------------------------------- +# Builtin tools +# --------------------------------------------------------------------------- + +@tools.router.get("/builtin") +async def list_builtin_tools() -> dict: + return {"tools": BUILTIN_TOOLS} + + +def load_builtin_permissions() -> dict[str, str]: + if not os.path.exists(BUILTIN_PERMS_PATH): + return {} + with open(BUILTIN_PERMS_PATH) as f: + return json.load(f) + + +def save_builtin_permissions(perms: dict[str, str]) -> None: + os.makedirs(os.path.dirname(BUILTIN_PERMS_PATH), exist_ok=True) + with open(BUILTIN_PERMS_PATH, "w") as f: + json.dump(perms, f, indent=2) + + +@tools.router.get("/builtin/permissions") +async def get_builtin_permissions() -> dict: + return {"permissions": load_builtin_permissions()} + + +@tools.router.put("/builtin/permissions") +async def update_builtin_permissions(body: dict) -> dict: + valid_names = {t["name"] for t in BUILTIN_TOOLS} + valid_policies = {"allow", "ask", "deny"} + perms = load_builtin_permissions() + for name, policy in body.get("permissions", {}).items(): + if name in valid_names and policy in valid_policies: + perms[name] = policy + save_builtin_permissions(perms) + return {"permissions": perms} + + +# --------------------------------------------------------------------------- +# User-installed tool CRUD +# --------------------------------------------------------------------------- + +@tools.router.get("/list") +async def list_tools() -> dict: + return {"tools": [t.model_dump() for t in TOOL_STORE.load_all()]} + + +class ToolCreate(BaseModel): + name: str + description: str = "" + command: str = "" + mcp_config: dict[str, Any] = Field(default_factory=dict) + credentials: dict[str, str] = Field(default_factory=dict) + auth_type: str = "none" + auth_status: str = "none" + oauth_provider: Optional[str] = None + +@tools.router.post("/create") +async def create_tool(body: ToolCreate) -> dict: + tool = ToolDefinition( + name=body.name, + description=body.description, + command=body.command, + mcp_config=body.mcp_config, + credentials=body.credentials, + auth_type=body.auth_type, + auth_status=body.auth_status, + oauth_provider=body.oauth_provider, + ) + TOOL_STORE.save(tool) + return {"ok": True, "tool": tool.model_dump()} + + +@tools.router.get("/{tool_id}") +async def get_tool(tool_id: str) -> dict: + return TOOL_STORE.load(tool_id).model_dump() + + +class ToolUpdate(BaseModel): + name: Optional[str] = None + description: Optional[str] = None + command: Optional[str] = None + mcp_config: Optional[dict[str, Any]] = None + credentials: Optional[dict[str, str]] = None + auth_type: Optional[str] = None + auth_status: Optional[str] = None + oauth_provider: Optional[str] = None + oauth_tokens: Optional[dict[str, Any]] = None + tool_permissions: Optional[dict[str, TOOL_PERMISSIONS]] = None + connected_account_email: Optional[str] = None + enabled: Optional[bool] = None + +@tools.router.put("/{tool_id}") +async def update_tool(tool_id: str, body: ToolUpdate) -> dict: + tool = TOOL_STORE.load(tool_id) + for k, v in body.model_dump(exclude_none=True).items(): + setattr(tool, k, v) + TOOL_STORE.save(tool) + return {"ok": True, "tool": tool.model_dump()} + + +@tools.router.delete("/{tool_id}") +async def delete_tool(tool_id: str) -> dict: + TOOL_STORE.delete(tool_id) + return {"ok": True} + + +# --------------------------------------------------------------------------- +# Discovery +# --------------------------------------------------------------------------- + +@tools.router.post("/{tool_id}/discover") +async def discover(tool_id: str) -> dict: + tool = TOOL_STORE.load(tool_id) + + if tool.auth_type == "oauth2" and tool.auth_status == "connected": + refreshed = await refresh_oauth_token(tool) + if not refreshed and tool.oauth_tokens.get("access_token"): + expiry = tool.oauth_tokens.get("token_expiry", 0) + if isinstance(expiry, (int, float)) and time.time() >= expiry - 60: + raise HTTPException( + status_code=502, + detail="OAuth token expired and refresh failed. Try reconnecting.", + ) + + mcp_tool = tool_definition_to_mcp_tool(tool, oauth_providers=OAUTH_PROVIDERS) + if not mcp_tool: + raise HTTPException(status_code=400, detail="Cannot derive MCP config for tool") + + config = list[McpServerConfig](mcp_tool.to_mcp_server_config().values())[0] + if isinstance(config, dict): + discovery_config = config + else: + raise HTTPException(status_code=400, detail="Unexpected MCP config format") + + try: + raw_tools = await discover_tools(discovery_config, tool_name=tool.name) + except DiscoveryConfigError as e: + raise HTTPException(status_code=400, detail=str(e)) + except DiscoveryError as e: + raise HTTPException(status_code=502, detail=str(e)) + except Exception as e: + msg = str(e).strip() or type(e).__name__ + logger.warning(f"MCP tool discovery failed for {tool.name}: {msg}", exc_info=True) + raise HTTPException(status_code=502, detail=f"Discovery failed: {msg}") + + permissions: dict[str, Any] = {} + for t in raw_tools: + name = t["name"] + permissions[name] = tool.tool_permissions.get(name, "ask") + + permissions["_tool_descriptions"] = {t["name"]: t["description"] for t in raw_tools} + permissions["_tool_schemas"] = { + t["name"]: t.get("inputSchema") for t in raw_tools if t.get("inputSchema") + } + + tool.tool_permissions = permissions + TOOL_STORE.save(tool) + + return {"ok": True, "tool": tool.model_dump()} + + +@tools.router.get("/load_user_toolkit") +async def load_user_toolkit() -> Optional[Toolkit]: + """Load all user-installed tools from the store and return them as a Toolkit. + + Returns None if no valid tools could be converted. + """ + mcp_tools: list[Tool] = [] + for td in TOOL_STORE.load_all(): + if not td.mcp_config or not td.enabled: + continue + if td.auth_status not in ("configured", "connected", "none"): + continue + mcp_tool: Optional[MCP_Tool] = tool_definition_to_mcp_tool(td, oauth_providers=OAUTH_PROVIDERS) + if mcp_tool is None: + continue + if td.tool_permissions: + known = set[str](td.tool_descriptions.keys()) + if known: + denied = {k for k, v in td.tool_permissions.items() if v == "deny"} + if known <= denied: + mcp_tool.permission = "deny" + mcp_tools.append(mcp_tool) + + if not mcp_tools: + return None + + return Toolkit( + name="user_installed", + description="User-installed MCP tool servers", + tools=mcp_tools, + ) + +# --------------------------------------------------------------------------- +# OAuth routes +# --------------------------------------------------------------------------- + +tools.router.add_api_route("/oauth/callback", oauth.oauth_callback, methods=["GET"]) +tools.router.add_api_route("/{tool_id}/oauth/start", oauth.oauth_start, methods=["POST"]) +tools.router.add_api_route("/{tool_id}/oauth/disconnect", oauth.oauth_disconnect, methods=["POST"]) diff --git a/backend/core/tools/shared_structs/MCP_Tool.py b/backend/core/tools/shared_structs/MCP_Tool.py index 6a9c484c..88498db9 100644 --- a/backend/core/tools/shared_structs/MCP_Tool.py +++ b/backend/core/tools/shared_structs/MCP_Tool.py @@ -16,7 +16,7 @@ from backend.core.tools.shared_structs.Tool import Tool class MCP_Tool(Tool): server_name: str - input_schema: type + input_schema: Optional[type] = None @typechecked def to_sdk_args(self) -> str: @@ -28,7 +28,7 @@ class MCP_Tool(Tool): class SDK_MCP_Tool(MCP_Tool): - # sdk transport: in-process handler + input_schema: type # required for SDK tools (overrides Optional on base) handler: Callable[[Dict[str, Any]], Awaitable[Dict[str, Any]]] @typechecked diff --git a/backend/main.py b/backend/main.py index 823be5e8..04e2380f 100644 --- a/backend/main.py +++ b/backend/main.py @@ -12,10 +12,11 @@ from backend.apps.dashboards.dashboards import dashboards from backend.apps.health.health import health from backend.apps.settings.settings import settings from backend.apps.modes.modes import modes +from backend.apps.tools.tools import tools from fastapi.middleware.cors import CORSMiddleware main_app = MainApp([ - health, agents, settings, dashboards, modes + health, tools, agents, settings, dashboards, modes ]) app = main_app.app