From ed73ef3d8df988b0e9646b82df5b1a453202ef8d Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 19 May 2026 21:35:04 +0400 Subject: [PATCH] refac --- backend/open_webui/utils/headers.py | 3 + backend/open_webui/utils/middleware.py | 150 +++++++++++-------------- backend/open_webui/utils/tools.py | 103 +++++++++++------ 3 files changed, 138 insertions(+), 118 deletions(-) diff --git a/backend/open_webui/utils/headers.py b/backend/open_webui/utils/headers.py index 510f60b335..1d616627f7 100644 --- a/backend/open_webui/utils/headers.py +++ b/backend/open_webui/utils/headers.py @@ -28,6 +28,8 @@ def get_custom_headers(custom_headers: dict, user=None, metadata: dict = None) - '{{MESSAGE_ID}}': metadata.get('message_id', '') or '', '{{USER_ID}}': (user.id if user else '') or '', '{{USER_NAME}}': (user.name if user else '') or '', + '{{USER_EMAIL}}': (user.email if user else '') or '', + '{{USER_ROLE}}': (user.role if user else '') or '', } parsed_headers = {} @@ -39,3 +41,4 @@ def get_custom_headers(custom_headers: dict, user=None, metadata: dict = None) - parsed_headers[key] = value return parsed_headers + diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index a6a9697c18..b78ad9f701 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -33,12 +33,9 @@ from open_webui.env import ( CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES, CHAT_RESPONSE_STREAM_DELTA_CHUNK_SIZE, ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION, - ENABLE_FORWARD_USER_INFO_HEADERS, ENABLE_QUERIES_CACHE, ENABLE_REALTIME_CHAT_SAVE, ENABLE_RESPONSES_API_STATEFUL, - FORWARD_SESSION_INFO_HEADER_CHAT_ID, - FORWARD_SESSION_INFO_HEADER_MESSAGE_ID, GLOBAL_LOG_LEVEL, RAG_SYSTEM_CONTEXT, ) @@ -89,7 +86,7 @@ from open_webui.utils.filter import ( get_sorted_filter_ids, process_filter_functions, ) -from open_webui.utils.headers import include_user_info_headers + from open_webui.utils.mcp.client import MCPClient from open_webui.utils.misc import ( add_or_update_system_message, @@ -121,6 +118,7 @@ from open_webui.utils.task import ( tools_function_calling_generation_template, ) from open_webui.utils.tools import ( + build_tool_server_headers, get_builtin_tools, get_terminal_tools, get_tools, @@ -2237,6 +2235,62 @@ def strip_skill_mentions(messages: list[dict]) -> None: part['text'] = strip_re.sub('', text).strip() + +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. + """ + mcp_server_connection = None + for server_connection in request.app.state.config.TOOL_SERVER_CONNECTIONS: + if ( + server_connection.get('type', '') == 'mcp' + and server_connection.get('info', {}).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): # Pipeline Inlet -> Filter Inlet -> Chat Memory -> Chat Web Search -> Chat Image Generation # -> Chat Code Interpreter (Form Data Update) -> (Default) Chat Tools Function Calling @@ -2640,83 +2694,18 @@ async def process_chat_payload(request, form_data, user, metadata, model): for tool_id in tool_ids: if tool_id.startswith('server:mcp:'): try: - server_id = tool_id[len('server:mcp:') :] + server_id = tool_id[len('server:mcp:'):] - mcp_server_connection = None - for server_connection in request.app.state.config.TOOL_SERVER_CONNECTIONS: - if ( - server_connection.get('type', '') == 'mcp' - and server_connection.get('info', {}).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') + result = await connect_mcp_server( + request, server_id, user, metadata, extra_params, + ) + if result is None: continue - # Check access control for MCP server - if not await has_connection_access(user, mcp_server_connection): - log.warning(f'Access denied to MCP server {server_id} for user {user.id}') - continue + client, tool_specs = result + mcp_clients[server_id] = client - auth_type = mcp_server_connection.get('auth_type', '') - headers = {} - if auth_type == 'bearer': - headers['Authorization'] = f'Bearer {mcp_server_connection.get("key", "")}' - elif auth_type == 'none': - # No authentication - pass - elif auth_type == 'session': - headers['Authorization'] = f'Bearer {request.state.token.credentials}' - elif auth_type == 'system_oauth': - oauth_token = extra_params.get('__oauth_token__', None) - if oauth_token: - headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' - elif auth_type in ('oauth_2.1', 'oauth_2.1_static'): - try: - splits = server_id.split(':') - server_id = splits[-1] if len(splits) > 1 else server_id - - oauth_token = await request.app.state.oauth_client_manager.get_oauth_token( - user.id, f'mcp:{server_id}' - ) - - if oauth_token: - headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' - except Exception as e: - log.error(f'Error getting OAuth token: {e}') - oauth_token = None - - connection_headers = mcp_server_connection.get('headers', None) - if connection_headers and isinstance(connection_headers, dict): - for key, value in connection_headers.items(): - headers[key] = value - - # Add user info headers if enabled - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - if metadata and metadata.get('chat_id'): - headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id') - if metadata and metadata.get('message_id'): - headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id') - - mcp_clients[server_id] = MCPClient() - await mcp_clients[server_id].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 mcp_clients[server_id].list_tool_specs() for tool_spec in tool_specs: - async def make_tool_function(client, function_name): async def tool_function(**kwargs): return await client.call_tool( @@ -2726,12 +2715,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): return tool_function - if function_name_filter_list: - if not is_string_allowed(tool_spec['name'], function_name_filter_list): - # Skip this function - continue - - tool_function = await make_tool_function(mcp_clients[server_id], tool_spec['name']) + tool_function = await make_tool_function(client, tool_spec['name']) mcp_tools_dict[f'{server_id}_{tool_spec["name"]}'] = { 'spec': { @@ -2740,7 +2724,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): }, 'callable': tool_function, 'type': 'mcp', - 'client': mcp_clients[server_id], + 'client': client, 'direct': False, } except Exception as e: diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 4017d0c48d..23dfa6b607 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -101,6 +101,68 @@ from pydantic.fields import FieldInfo log = logging.getLogger(__name__) +async def build_tool_server_headers( + connection: dict, + request, + user, + server_id: str = '', + metadata: dict | None = None, + extra_params: dict | None = None, +) -> tuple[dict, dict]: + """Build auth headers and cookies for a tool server connection. + + Handles bearer, session, system_oauth, and oauth_2.1 auth types plus + custom header interpolation and user-info forwarding. + Shared by MCP and OpenAPI paths. + + Returns (headers, cookies). + """ + extra_params = extra_params or {} + metadata = metadata or {} + + auth_type = connection.get('auth_type', 'bearer') + headers = {} + cookies = {} + + if auth_type == 'bearer': + headers['Authorization'] = f'Bearer {connection.get("key", "")}' + elif auth_type == 'session': + cookies = request.cookies if hasattr(request, 'cookies') else {} + headers['Authorization'] = f'Bearer {request.state.token.credentials}' + elif auth_type == 'system_oauth': + cookies = request.cookies if hasattr(request, 'cookies') else {} + oauth_token = extra_params.get('__oauth_token__', None) + if oauth_token: + headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' + elif auth_type in ('oauth_2.1', 'oauth_2.1_static'): + try: + splits = server_id.split(':') + oauth_server_id = splits[-1] if len(splits) > 1 else server_id + connection_type = connection.get('type', 'openapi') + oauth_token = await request.app.state.oauth_client_manager.get_oauth_token( + user.id, f'{connection_type}:{oauth_server_id}' + ) + if oauth_token: + headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' + except Exception as e: + log.error(f'Error getting OAuth token: {e}') + + # Interpolate template vars in custom connection headers + connection_headers = connection.get('headers', None) + if connection_headers and isinstance(connection_headers, dict): + headers.update(get_custom_headers(connection_headers, user, metadata)) + + # Add user info headers if enabled + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) + if metadata.get('chat_id'): + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata['chat_id'] + if metadata.get('message_id'): + headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata['message_id'] + + return headers, cookies + + # 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( @@ -312,41 +374,12 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # Skip this function continue - auth_type = tool_server_connection.get('auth_type', 'bearer') - - cookies = {} - headers = { - 'Content-Type': 'application/json', - } - - if auth_type == 'bearer': - headers['Authorization'] = f'Bearer {tool_server_connection.get("key", "")}' - elif auth_type == 'none': - # No authentication - pass - elif auth_type == 'session': - cookies = request.cookies - headers['Authorization'] = f'Bearer {request.state.token.credentials}' - elif auth_type == 'system_oauth': - cookies = request.cookies - oauth_token = extra_params.get('__oauth_token__', None) - if oauth_token: - headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' - - connection_headers = tool_server_connection.get('headers', None) - if connection_headers and isinstance(connection_headers, dict): - metadata = extra_params.get('__metadata__', {}) - custom_headers = get_custom_headers(connection_headers, user, metadata) - headers.update(custom_headers) - - # Add user info headers if enabled - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - metadata = extra_params.get('__metadata__', {}) - if metadata and metadata.get('chat_id'): - headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id') - if metadata and metadata.get('message_id'): - headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id') + metadata = extra_params.get('__metadata__', {}) + headers, cookies = await build_tool_server_headers( + tool_server_connection, request, user, + server_id=server_id, metadata=metadata, extra_params=extra_params, + ) + headers.setdefault('Content-Type', 'application/json') async def make_tool_function(function_name, tool_server_data, headers): async def tool_function(**kwargs):