This commit is contained in:
Timothy Jaeryang Baek
2026-05-19 21:35:04 +04:00
parent f02aeea0bb
commit ed73ef3d8d
3 changed files with 138 additions and 118 deletions
+3
View File
@@ -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
+67 -83
View File
@@ -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:
+68 -35
View File
@@ -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):