This commit is contained in:
Timothy Jaeryang Baek 2026-10-08 20:33:47 +04:00
parent 87a9374595
commit 93fc3fcb72
7 changed files with 85 additions and 95 deletions

View file

@ -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__,

View file

@ -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

View file

@ -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,

View file

@ -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']

View file

@ -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 = {

View file

@ -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

View file

@ -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,