diff --git a/backend/open_webui/functions.py b/backend/open_webui/functions.py index a0ed278af2..ed9e31a242 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -36,6 +36,7 @@ from open_webui.utils.misc import ( openai_chat_completion_message_template, prepend_to_first_user_message_content, ) +from open_webui.utils.oauth import get_system_oauth_token from open_webui.utils.payload import ( apply_model_params_to_body_openai, apply_system_prompt_to_body, @@ -242,28 +243,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di __task__ = metadata.get('task', None) __task_body__ = metadata.get('task_body', None) - oauth_token = None - try: - oauth_session_id = request.cookies.get('oauth_session_id', None) - if oauth_session_id: - oauth_token = await request.app.state.oauth_manager.get_oauth_token( - user.id, - oauth_session_id, - ) - - # Fallback: no cookie (automation, API key, etc.) — use most recent session - if oauth_token is None: - from open_webui.models.oauth_sessions import OAuthSessions - - sessions = await OAuthSessions.get_sessions_by_user_id(user.id) - if sessions: - best = max(sessions, key=lambda s: s.updated_at) - oauth_token = await request.app.state.oauth_manager.get_oauth_token( - user.id, - best.id, - ) - except Exception as e: - log.error(f'Error getting OAuth token: {e}') + oauth_token = await get_system_oauth_token(request, user) extra_params = { '__event_emitter__': __event_emitter__, diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index ecd48f4506..d5bf49dfef 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -209,7 +209,7 @@ async def get_tool_specs(request: Request, id: str, user=Depends(get_verified_us 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, {}, {}) + 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 diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index 35256b89b2..49dec5faee 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -57,17 +57,41 @@ def create_insecure_httpx_client(headers=None, timeout=None, auth=None): return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=False) +class OAuthTokenAuth(httpx.Auth): + """Resolve current credentials per request and recover from concurrent token rotation.""" + + requires_request_body = True + + def __init__(self, get_headers): + self.get_headers = get_headers + + async def async_auth_flow(self, request): + headers = httpx.Headers(await self.get_headers()) + authorization = headers.get('Authorization') + if not authorization: + raise httpx.RequestError('No OAuth access token available', request=request) + request.headers.update(headers) + response = yield request + + if response.status_code == 401: + headers = httpx.Headers(await self.get_headers()) + if headers.get('Authorization') and headers['Authorization'] != authorization: + request.headers.update(headers) + yield request + + class MCPClient: def __init__(self): self.session: Optional[ClientSession] = None self.exit_stack = None - async def connect(self, url: str, headers: Optional[dict] = None): + async def connect(self, url: str, headers: Optional[dict] = None, auth: Optional[httpx.Auth] = None): async with AsyncExitStack() as exit_stack: try: self._streams_context = streamablehttp_client( url, headers=headers, + auth=auth, httpx_client_factory=create_httpx_client if AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL else create_insecure_httpx_client, diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 08182cab67..537688ba72 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -53,7 +53,6 @@ from open_webui.models.config import Config from open_webui.models.folders import Folders from open_webui.models.models import Models from open_webui.models.notes import Notes -from open_webui.models.oauth_sessions import OAuthSessions from open_webui.models.users import UserModel, Users from open_webui.retrieval.utils import filter_source_metadata, get_sources_from_items from open_webui.routers.images import ( @@ -127,6 +126,7 @@ from open_webui.utils.misc import ( set_last_user_message_content, strip_empty_content_blocks, ) +from open_webui.utils.oauth import get_system_oauth_token from open_webui.utils.payload import apply_params_to_form_data, apply_system_prompt_to_body, resolve_system_prompt from open_webui.utils.plugin import load_function_module_by_id from open_webui.utils.response import merge_usage, normalize_usage @@ -2952,7 +2952,6 @@ async def process_chat_payload(request, form_data, user, metadata, model): server_id, user, metadata, - extra_params, ) if result is None: continue @@ -3741,42 +3740,6 @@ def update_assistant_message_from_stream(assistant_message, raw): assistant_message['content'] = '' + content -async def get_system_oauth_token(request, user): - """Get the system OAuth token for a user. - - Primary path: use the oauth_session_id cookie (browser requests). - Fallback: look up the user's most recent OAuth session from the DB - (covers automations, API calls, and other cookie-less contexts). - """ - oauth_token = None - try: - oauth_session_id = request.cookies.get('oauth_session_id', None) - if oauth_session_id: - oauth_token = await request.app.state.oauth_manager.get_oauth_token( - user.id, - oauth_session_id, - ) - - # Fallback: no cookie (automation, API key, etc.) — use most recent session - if oauth_token is None: - from open_webui.models.oauth_sessions import OAuthSessions - - sessions = await OAuthSessions.get_sessions_by_user_id(user.id) - # Filter out MCP-provider sessions — their token refresh is handled - # separately by oauth_client_manager. Passing them to the SSO - # oauth_manager causes a failed refresh and session deletion (#24618). - sessions = [s for s in sessions if not (s.provider or '').startswith('mcp:')] - if sessions: - best = max(sessions, key=lambda s: s.updated_at) - oauth_token = await request.app.state.oauth_manager.get_oauth_token( - user.id, - best.id, - ) - except Exception as e: - log.error(f'Error getting OAuth token: {e}') - return oauth_token - - async def background_tasks_handler(ctx): request = ctx['request'] form_data = ctx['form_data'] diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 4de0379cc2..04a4831d19 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -127,6 +127,41 @@ from open_webui.utils.json_codec import JSONCodec logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) log = logging.getLogger(__name__) + +async def get_system_oauth_token(request, user): + """Get the system OAuth token for a user. + + Primary path: use the oauth_session_id cookie (browser requests). + Fallback: look up the user's most recent OAuth session from the DB + (covers automations, API calls, and other cookie-less contexts). + """ + oauth_token = None + try: + oauth_session_id = request.cookies.get('oauth_session_id', None) + if oauth_session_id: + oauth_token = await request.app.state.oauth_manager.get_oauth_token( + user.id, + oauth_session_id, + ) + + # Fallback: no cookie (automation, API key, etc.) — use most recent session + if oauth_token is None: + sessions = await OAuthSessions.get_sessions_by_user_id(user.id) + # Filter out MCP-provider sessions — their token refresh is handled + # separately by oauth_client_manager. Passing them to the SSO + # oauth_manager causes a failed refresh and session deletion (#24618). + sessions = [s for s in sessions if not (s.provider or '').startswith('mcp:')] + if sessions: + best = max(sessions, key=lambda s: s.updated_at) + oauth_token = await request.app.state.oauth_manager.get_oauth_token( + user.id, + best.id, + ) + except Exception as e: + log.error(f'Error getting OAuth token: {e}') + return oauth_token + + OAUTH_RESOURCE_PARAMETER_MODES = {'auto', 'include', 'omit'} OAUTH_RUNTIME_CONFIG = { diff --git a/backend/open_webui/utils/terminals.py b/backend/open_webui/utils/terminals.py index dd6897e01c..c764ca2efb 100644 --- a/backend/open_webui/utils/terminals.py +++ b/backend/open_webui/utils/terminals.py @@ -164,7 +164,6 @@ async def get_terminal_json(request, user, metadata: dict, path: str, extra_para request, user_model, metadata=metadata, - extra_params=extra_params, ) headers['Accept'] = 'application/json' headers['X-User-Id'] = user_model.id diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 510dc71de3..8f22cfc3a0 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -110,8 +110,9 @@ 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.mcp.client import MCPClient, OAuthTokenAuth from open_webui.utils.misc import is_string_allowed +from open_webui.utils.oauth import get_system_oauth_token from open_webui.utils.plugin import get_tool_contents_cache, get_tools_cache, load_tool_module_by_id from open_webui.utils.terminals import ( TERMINAL_CONTEXT_HEADER, @@ -132,7 +133,6 @@ async def build_tool_server_headers( 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. @@ -142,7 +142,6 @@ async def build_tool_server_headers( Returns (headers, cookies). """ - extra_params = extra_params or {} metadata = metadata or {} auth_type = connection.get('auth_type', 'bearer') @@ -154,7 +153,7 @@ async def build_tool_server_headers( elif auth_type == 'session': headers.update(bearer_auth_header(request.state.token.credentials)) elif auth_type == 'system_oauth': - oauth_token = extra_params.get('__oauth_token__', None) + oauth_token = await get_system_oauth_token(request, user) if oauth_token: headers.update(bearer_auth_header(oauth_token.get('access_token', ''))) elif auth_type in ('oauth_2.1', 'oauth_2.1_static'): @@ -173,7 +172,8 @@ async def build_tool_server_headers( # Interpolate template vars in custom connection headers connection_headers = connection.get('headers', None) if connection_headers and isinstance(connection_headers, dict): - headers.update(await get_custom_headers(connection_headers, user, metadata)) + for key, value in (await get_custom_headers(connection_headers, user, metadata)).items(): + headers['Authorization' if key.lower() == 'authorization' else key] = value # Add user info headers if enabled if ENABLE_FORWARD_USER_INFO_HEADERS and user: @@ -191,7 +191,6 @@ async def connect_mcp_server( 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). @@ -215,22 +214,14 @@ async def connect_mcp_server( 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), - } + async def get_headers(): + headers, _ = await build_tool_server_headers( + mcp_server_connection, request, user, server_id=server_id, metadata=metadata + ) + return headers - headers, _ = await build_tool_server_headers( - mcp_server_connection, - request, - user, - server_id=server_id, - metadata=metadata, - extra_params=extra_params, - ) + auth = OAuthTokenAuth(get_headers) if mcp_server_connection.get('auth_type') == 'system_oauth' else None + headers = {} if auth else await get_headers() 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') @@ -240,6 +231,7 @@ async def connect_mcp_server( await client.connect( url=mcp_server_connection.get('url', ''), headers=headers if headers else None, + auth=auth, ) function_name_filter_list = (mcp_server_connection.get('config') or {}).get('function_name_filter_list', '') if isinstance(function_name_filter_list, str): @@ -504,18 +496,13 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr continue 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, cookies): + async def make_tool_function(function_name, tool_server_data, connection, server_id, metadata): async def tool_function(**kwargs): + headers, cookies = await build_tool_server_headers( + connection, request, user, server_id=server_id, metadata=metadata + ) + headers.setdefault('Content-Type', 'application/json') return await execute_tool_server( url=tool_server_data['url'], headers=headers, @@ -527,7 +514,9 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tool_function - tool_function = await make_tool_function(function_name, tool_server_data, headers, cookies) + tool_function = await make_tool_function( + function_name, tool_server_data, tool_server_connection, server_id, metadata + ) callable = await get_async_tool_function_and_apply_extra_params( tool_function,