refac
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user