From fb741ebcd2daed626405f9458c192e5161bce4c0 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 5 Oct 2026 10:56:06 +0400 Subject: [PATCH] refac --- backend/open_webui/routers/tools.py | 34 +++- backend/open_webui/utils/middleware.py | 61 +------ backend/open_webui/utils/tools.py | 76 ++++++++- src/lib/apis/tools/index.ts | 10 ++ src/lib/components/chat/MessageInput.svelte | 7 +- .../components/chat/ToolServersModal.svelte | 154 +++++++++++++----- src/lib/i18n/locales/en-US/translation.json | 5 + 7 files changed, 240 insertions(+), 107 deletions(-) diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 2191f82db0..6d1ff070b3 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import logging import re import time @@ -44,7 +45,8 @@ from open_webui.utils.plugin import ( replace_imports, resolve_valves_schema_options, ) -from open_webui.utils.tools import get_tool_servers, get_tool_specs +from open_webui.utils.tools import connect_mcp_server, get_tool_servers +from open_webui.utils.tools import get_tool_specs as get_local_tool_specs from pydantic import BaseModel, HttpUrl from sqlalchemy.ext.asyncio import AsyncSession @@ -196,6 +198,32 @@ async def get_tools( return tools +@router.get('/id/{id}/specs') +async def get_tool_specs(request: Request, id: str, user=Depends(get_verified_user)): + """Discover tools for an accessible connection. Currently supports MCP.""" + if not id.startswith('server:mcp:'): + raise HTTPException(status_code=404, detail='Tool not found') + + try: + # Keep connect, discovery and cleanup in one task for the MCP transport. + async with asyncio.timeout(15): + result = await connect_mcp_server(request, id.removeprefix('server:mcp:'), user, {}, {}) + if result is None: + raise HTTPException(status_code=404, detail='Tool not found') + client, specs = result + try: + return {'specs': [{'name': spec['name'], 'description': spec.get('description', '')} for spec in specs]} + finally: + await client.disconnect() + except HTTPException: + raise + except TimeoutError: + raise HTTPException(status_code=504, detail='Tool discovery timed out') + except Exception: + log.exception('Failed to discover tool specs') + raise HTTPException(status_code=502, detail='Unable to load tools') + + ############################ # GetToolList ############################ @@ -397,7 +425,7 @@ async def create_new_tools( TOOLS = get_tools_cache(request) TOOLS[form_data.id] = tool_module - specs = get_tool_specs(TOOLS[form_data.id]) + specs = get_local_tool_specs(TOOLS[form_data.id]) tools = await Tools.insert_new_tool(user.id, form_data, specs, db=db) tool_cache_dir = CACHE_DIR / 'tools' / form_data.id @@ -539,7 +567,7 @@ async def update_tools_by_id( TOOLS = get_tools_cache(request) TOOLS[id] = tool_module - specs = get_tool_specs(TOOLS[id]) + specs = get_local_tool_specs(TOOLS[id]) form_data.access_grants = await filter_allowed_access_grants( await Config.get('user.permissions'), diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index c6c0ed83c1..a48f42d960 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -83,7 +83,7 @@ from open_webui.socket.main import ( get_event_emitter, ) from open_webui.tasks import clear_response_stream, save_response_stream -from open_webui.utils.access_control import has_connection_access, has_permission +from open_webui.utils.access_control import has_permission from open_webui.utils.access_control.files import get_owner_accessible_folder_files from open_webui.utils.access_control.folders import has_folder_access from open_webui.utils.ask_user import stage_ask_user_tool_calls @@ -104,7 +104,6 @@ from open_webui.utils.filter import ( process_filter_functions, ) from open_webui.utils.json_codec import JSONCodec -from open_webui.utils.mcp.client import MCPClient from open_webui.utils.memory import add_memory_context, review_memory_after_turn from open_webui.utils.misc import ( add_or_update_system_message, @@ -122,7 +121,6 @@ from open_webui.utils.misc import ( get_response_error_detail, get_system_message, is_raster_image_content_type, - is_string_allowed, merge_system_messages, prepend_to_first_user_message_content, replace_system_message_content, @@ -145,7 +143,7 @@ from open_webui.utils.task import ( tools_function_calling_generation_template, ) from open_webui.utils.tools import ( - build_tool_server_headers, + connect_mcp_server, get_attached_knowledge, get_builtin_tools, get_terminal_tools, @@ -2320,61 +2318,6 @@ def sanitize_tool_pairs(messages: list[dict]) -> list[dict]: return sanitized -async def connect_mcp_server( - request, - server_id: str, - user, - metadata: dict, - extra_params: dict, -) -> tuple[MCPClient, list[dict]] | None: - """Resolve an MCP server connection, authenticate, and return (client, tool_specs). - - Returns None if the server is not found or access is denied. - """ - if not ENABLE_TOOL_SERVERS: - log.debug('MCP resolution skipped: external plugins are disabled') - return None - - mcp_server_connection = None - for server_connection in await Config.get('tool_server.connections', []): - if server_connection.get('type', '') == 'mcp' and (server_connection.get('info') or {}).get('id') == server_id: - mcp_server_connection = server_connection - break - - if not mcp_server_connection: - log.error(f'MCP server with id {server_id} not found') - return None - - if not await has_connection_access(user, mcp_server_connection): - log.warning(f'Access denied to MCP server {server_id} for user {user.id}') - return None - - headers, _ = await build_tool_server_headers( - mcp_server_connection, - request, - user, - server_id=server_id, - metadata=metadata, - extra_params=extra_params, - ) - - client = MCPClient() - await client.connect( - url=mcp_server_connection.get('url', ''), - headers=headers if headers else None, - ) - - function_name_filter_list = mcp_server_connection.get('config', {}).get('function_name_filter_list', '') - if isinstance(function_name_filter_list, str): - function_name_filter_list = function_name_filter_list.split(',') - - tool_specs = await client.list_tool_specs() - if function_name_filter_list: - tool_specs = [spec for spec in tool_specs if is_string_allowed(spec['name'], function_name_filter_list)] - - return client, tool_specs - - async def process_chat_payload(request, form_data, user, metadata, model): # Ensure chat_id is always a string — external API clients may omit it. if not isinstance(metadata.get('chat_id'), str): diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 7edd7f401a..c00854c58a 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -21,7 +21,7 @@ from urllib.parse import quote, urlencode import aiohttp import yaml -from fastapi import Request +from fastapi import HTTPException, Request from langchain_core.utils.function_calling import ( convert_to_openai_function as convert_pydantic_model_to_openai_function_spec, ) @@ -110,6 +110,7 @@ from open_webui.utils.headers import ( normalize_bearer_token, ) from open_webui.utils.json_codec import JSONCodec +from open_webui.utils.mcp.client import MCPClient from open_webui.utils.misc import is_string_allowed from open_webui.utils.plugin import get_tool_contents_cache, get_tools_cache, load_tool_module_by_id from open_webui.utils.terminals import ( @@ -185,6 +186,75 @@ async def build_tool_server_headers( return headers, cookies +async def connect_mcp_server( + request, + server_id: str, + user, + metadata: dict, + extra_params: dict, +) -> tuple[MCPClient, list[dict]] | None: + """Resolve an MCP server connection, authenticate, and return (client, tool_specs). + + Returns None if the server is not found or access is denied. + """ + if not ENABLE_TOOL_SERVERS: + log.debug('MCP resolution skipped: external plugins are disabled') + return None + + mcp_server_connection = None + for server_connection in await Config.get('tool_server.connections', []): + if server_connection.get('type', '') == 'mcp' and (server_connection.get('info') or {}).get('id') == server_id: + mcp_server_connection = server_connection + break + + if not mcp_server_connection or not (mcp_server_connection.get('config') or {}).get('enable'): + log.error(f'MCP server with id {server_id} not found') + return None + + if not await has_connection_access(user, mcp_server_connection): + log.warning(f'Access denied to MCP server {server_id} for user {user.id}') + return None + + if mcp_server_connection.get('auth_type') == 'system_oauth' and not extra_params.get('__oauth_token__'): + session_id = request.cookies.get('oauth_session_id') + if session_id: + extra_params = { + **extra_params, + '__oauth_token__': await request.app.state.oauth_manager.get_oauth_token(user.id, session_id), + } + + headers, _ = await build_tool_server_headers( + mcp_server_connection, + request, + user, + server_id=server_id, + metadata=metadata, + extra_params=extra_params, + ) + + if mcp_server_connection.get('auth_type') in ('oauth_2.1', 'oauth_2.1_static') and not headers.get('Authorization'): + raise HTTPException(status_code=401, detail='Auth required') + + client = MCPClient() + try: + await client.connect( + url=mcp_server_connection.get('url', ''), + headers=headers if headers else None, + ) + function_name_filter_list = (mcp_server_connection.get('config') or {}).get('function_name_filter_list', '') + if isinstance(function_name_filter_list, str): + function_name_filter_list = function_name_filter_list.split(',') + + tool_specs = await client.list_tool_specs() + if function_name_filter_list: + tool_specs = [spec for spec in tool_specs if is_string_allowed(spec['name'], function_name_filter_list)] + return client, tool_specs + except BaseException: + # MCP sessions must be closed in the same task that opened them. + await client.disconnect() + raise + + # Let no function be called without need, and let what # it yields justify the cost of running it. async def get_async_tool_function_and_apply_extra_params( @@ -271,9 +341,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return {} enabled_ids = [ - tool_id - for tool_id in tool_ids - if (ENABLE_TOOL_SERVERS if tool_id.startswith('server:') else ENABLE_TOOLS) + tool_id for tool_id in tool_ids if (ENABLE_TOOL_SERVERS if tool_id.startswith('server:') else ENABLE_TOOLS) ] if len(enabled_ids) != len(tool_ids): log.debug('Excluded tools disabled by plugin configuration') diff --git a/src/lib/apis/tools/index.ts b/src/lib/apis/tools/index.ts index 8378299923..4f21efa5d9 100644 --- a/src/lib/apis/tools/index.ts +++ b/src/lib/apis/tools/index.ts @@ -99,6 +99,16 @@ export const getTools = async (token: string = '', query: string | null = null) return res; }; +export const getToolSpecs = async (token: string, id: string) => { + const res = await fetch(`${WEBUI_API_BASE_URL}/tools/id/${encodeURIComponent(id)}/specs`, { + headers: { Accept: 'application/json', authorization: `Bearer ${token}` } + }); + if (!res.ok) { + throw { status: res.status }; + } + return (await res.json()).specs; +}; + export const getToolList = async (token: string = '') => { let error = null; diff --git a/src/lib/components/chat/MessageInput.svelte b/src/lib/components/chat/MessageInput.svelte index 77cc07cc18..edb82e97f4 100644 --- a/src/lib/components/chat/MessageInput.svelte +++ b/src/lib/components/chat/MessageInput.svelte @@ -1640,7 +1640,12 @@ }); - + + oauthRedirectHandler({ id, serverId: id.split(':').at(-1), authType: 'mcp' }, chatInputDraft)} +/> + import { getToolSpecs } from '$lib/apis/tools'; import { resolveLocalizedResource } from '$lib/utils/localizedContent'; import { getContext } from 'svelte'; import { toolServers, tools } from '$lib/stores'; @@ -9,26 +10,62 @@ import XMark from '$lib/components/icons/XMark.svelte'; export let show = false; - export let selectedToolIds = []; + export let selectedToolIds: string[] = []; + export let onConnect: (id: string) => void = () => {}; - let selectedTools = []; + let discovery: Record< + string, + { specs?: any[]; loading?: boolean; error?: 'auth' | 'unavailable' } + > = {}; - $: selectedTools = ($tools ?? []).filter((tool) => selectedToolIds.includes(tool.id)); + const loadSpecs = async (tool: { id: string }) => { + if (discovery[tool.id]?.loading) return; + discovery = { ...discovery, [tool.id]: { loading: true } }; + try { + const specs = await getToolSpecs(localStorage.token, tool.id); + discovery = { ...discovery, [tool.id]: { specs } }; + } catch (error) { + discovery = { + ...discovery, + [tool.id]: { + error: + error && typeof error === 'object' && 'status' in error && error.status === 401 + ? 'auth' + : 'unavailable' + } + }; + } + }; - const i18n = getContext('i18n'); + const reconnect = (tool: { id: string }) => { + show = false; + onConnect(tool.id); + }; - const authStatus = (tool) => + let selectedTools: any[] = []; + + $: selectedTools = (($tools ?? []) as any[]).filter((tool) => selectedToolIds.includes(tool.id)); + + $: selectedToolServers = (($toolServers ?? []) as any[]).filter((server, idx) => + selectedToolIds.some((id) => { + if (!id.startsWith('direct_server:')) return false; + const serverId = id.slice('direct_server:'.length); + return !isNaN(parseInt(serverId)) ? parseInt(serverId) === idx : serverId === server?.id; + }) + ); + + const i18n = getContext('i18n'); + + const authStatus = (tool: { id: string; authenticated?: boolean }) => tool?.authenticated === false ? { label: $i18n.t('Auth required'), - dot: 'bg-amber-500', - pill: 'text-amber-700 dark:text-amber-300' + dot: 'bg-amber-500' } : tool?.authenticated === true ? { label: $i18n.t('Connected'), - dot: 'bg-green-500', - pill: 'text-green-700 dark:text-green-300' + dot: 'bg-green-500' } : null; @@ -49,7 +86,7 @@ {#if selectedTools.length > 0} - {#if $toolServers.length > 0} + {#if selectedToolServers.length > 0}
{$i18n.t('Tools')}
@@ -57,29 +94,24 @@
- {#each selectedTools as tool} - {@const status = authStatus(tool)} - {@const toolSpecs = tool?.specs ?? []} + {#each selectedTools as tool (tool.id)} + {@const isMcp = tool.id.startsWith('server:mcp:')} + {@const state = discovery[tool.id]} + {@const needsAuth = tool.authenticated === false || state?.error === 'auth'} + {@const status = authStatus(needsAuth ? { ...tool, authenticated: false } : tool)} + {@const toolSpecs = needsAuth ? undefined : isMcp ? state?.specs : tool.specs} 0} - disabled={toolSpecs.length === 0} + chevron + onChange={(open: boolean) => { + if (open && isMcp && !needsAuth && !state) loadSpecs(tool); + }} >
{resolveLocalizedResource(tool, $i18n.language)}
- {#if tool?.authenticated === false && status} - {status.label} - {/if} - {#if toolSpecs.length > 0} - - {toolSpecs.length} - - {/if} {#if status} {/if} + {#if needsAuth} + {$i18n.t('Auth required')} + {:else if toolSpecs !== undefined} + + {toolSpecs.length === 1 + ? $i18n.t('1 tool') + : $i18n.t('{{COUNT}} tools', { COUNT: toolSpecs.length })} + + {/if}
- {#if resolveLocalizedResource(tool, $i18n.language, 'description')}
{resolveLocalizedResource(tool, $i18n.language, 'description')}
{/if}
- -
- {#if toolSpecs.length > 0} - {#each toolSpecs as toolSpec} -
- {toolSpec?.name ?? toolSpec?.function?.name} -
- {/each} +
+ {#if needsAuth} + + {:else if state?.loading} +

{$i18n.t('Loading tools...')}

+ {:else if state?.error} +
+ {$i18n.t('Unable to load tools')} + +
+ {:else if toolSpecs?.length > 0} +
+ {#each toolSpecs as toolSpec} +
+
+ {toolSpec?.name ?? toolSpec?.function?.name} +
+ {#if toolSpec?.description ?? toolSpec?.function?.description} +
+ {toolSpec?.description ?? toolSpec?.function?.description} +
+ {/if} +
+ {/each} +
+ {:else} +

{$i18n.t('No tools found')}

{/if}
@@ -113,7 +182,7 @@
{/if} - {#if $toolServers.length > 0} + {#if selectedToolServers.length > 0}
{$i18n.t('Tool Servers')}
@@ -130,7 +199,7 @@ >
- {#each $toolServers as toolServer} + {#each selectedToolServers as toolServer}
@@ -146,14 +215,19 @@
-
+
{#each toolServer?.specs ?? [] as tool_spec} -
-
+
+
{tool_spec?.name}
-
+
{tool_spec?.description}
diff --git a/src/lib/i18n/locales/en-US/translation.json b/src/lib/i18n/locales/en-US/translation.json index cd1d3ebd55..018258aedf 100644 --- a/src/lib/i18n/locales/en-US/translation.json +++ b/src/lib/i18n/locales/en-US/translation.json @@ -1,4 +1,9 @@ { + "1 tool": "", + "{{COUNT}} tools": "", + "Reconnect": "", + "Loading tools...": "", + "Unable to load tools": "", "-1 for no limit, or a positive integer for a specific limit": "", "(latest)": "", "(leave blank for to use commercial endpoint)": "",