mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-09 03:18:18 +00:00
refac
This commit is contained in:
parent
87a9374595
commit
93fc3fcb72
7 changed files with 85 additions and 95 deletions
|
|
@ -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__,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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']
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue