mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-10 03:27:57 +00:00
refac
This commit is contained in:
parent
dd280082ac
commit
a0e606bfaa
1 changed files with 45 additions and 71 deletions
|
|
@ -499,6 +499,32 @@ async def get_discovery_urls(server_url) -> list[str]:
|
|||
return metadata.get_discovery_urls(server_url)
|
||||
|
||||
|
||||
async def _get_oauth_server_metadata(
|
||||
server_url: str, resource_metadata: ProtectedResourceMetadata
|
||||
) -> tuple[OAuthMetadata, str | None]:
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
for url in resource_metadata.get_discovery_urls(server_url):
|
||||
async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as response:
|
||||
if response.status != 200:
|
||||
continue
|
||||
try:
|
||||
return OAuthMetadata.model_validate(await response.json()), url
|
||||
except Exception as e:
|
||||
log.error(f'Error parsing OAuth metadata from {url}: {e}')
|
||||
|
||||
if resource_metadata.authorization_servers:
|
||||
raise ValueError(f'Could not discover the OAuth authorization server metadata for {server_url}')
|
||||
|
||||
# MCP 2025-03-26 defines origin-level defaults for servers without discovery.
|
||||
_, base_url = get_parsed_and_base_url(server_url)
|
||||
return OAuthMetadata(
|
||||
issuer=base_url,
|
||||
authorization_endpoint=f'{base_url}/authorize',
|
||||
token_endpoint=f'{base_url}/token',
|
||||
registration_endpoint=f'{base_url}/register',
|
||||
), None
|
||||
|
||||
|
||||
# TODO: Some OAuth providers require Initial Access Tokens (IATs) for dynamic client registration.
|
||||
# This is not currently supported.
|
||||
async def get_oauth_client_info_with_dynamic_client_registration(
|
||||
|
|
@ -509,9 +535,6 @@ async def get_oauth_client_info_with_dynamic_client_registration(
|
|||
oauth_scope: str | None = None,
|
||||
) -> OAuthClientInformationFull:
|
||||
try:
|
||||
oauth_server_metadata = None
|
||||
oauth_server_metadata_url = None
|
||||
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
|
||||
|
|
@ -538,53 +561,23 @@ async def get_oauth_client_info_with_dynamic_client_registration(
|
|||
elif resource_metadata.scopes_supported:
|
||||
oauth_client_metadata.scope = ' '.join(resource_metadata.scopes_supported)
|
||||
|
||||
discovery_urls = resource_metadata.get_discovery_urls(oauth_server_url)
|
||||
for url in discovery_urls:
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as oauth_server_metadata_response:
|
||||
if oauth_server_metadata_response.status == 200:
|
||||
try:
|
||||
oauth_server_metadata = OAuthMetadata.model_validate(
|
||||
await oauth_server_metadata_response.json()
|
||||
)
|
||||
oauth_server_metadata_url = url
|
||||
if (
|
||||
oauth_client_metadata.scope is None
|
||||
and oauth_server_metadata.scopes_supported is not None
|
||||
):
|
||||
oauth_client_metadata.scope = ' '.join(oauth_server_metadata.scopes_supported)
|
||||
oauth_server_metadata, oauth_server_metadata_url = await _get_oauth_server_metadata(
|
||||
oauth_server_url, resource_metadata
|
||||
)
|
||||
if oauth_client_metadata.scope is None and oauth_server_metadata.scopes_supported is not None:
|
||||
oauth_client_metadata.scope = ' '.join(oauth_server_metadata.scopes_supported)
|
||||
|
||||
if (
|
||||
oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
and oauth_client_metadata.token_endpoint_auth_method
|
||||
not in oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
):
|
||||
# Pick the first supported method from the server
|
||||
oauth_client_metadata.token_endpoint_auth_method = (
|
||||
oauth_server_metadata.token_endpoint_auth_methods_supported[0]
|
||||
)
|
||||
|
||||
break
|
||||
except Exception as e:
|
||||
log.error(f'Error parsing OAuth metadata from {url}: {e}')
|
||||
continue
|
||||
|
||||
# Fail fast if authorization server metadata discovery did not resolve an
|
||||
# authorization endpoint. Otherwise registration can still "succeed" (via
|
||||
# the /register fallback below) while issuer/server_metadata stay unset,
|
||||
# which later crashes at authorize time with authlib's
|
||||
# RuntimeError: Missing "authorize_url" value. (#26647)
|
||||
if oauth_server_metadata is None or not oauth_server_metadata.authorization_endpoint:
|
||||
log.error(f'OAuth authorization server metadata discovery failed for {oauth_server_url}')
|
||||
raise Exception(
|
||||
'Could not discover the OAuth authorization server metadata '
|
||||
f'(authorization_endpoint) for {oauth_server_url}. The MCP server must '
|
||||
'expose RFC 8414 / RFC 9728 discovery documents so Open WebUI can '
|
||||
'resolve where to send users to authorize.'
|
||||
if (
|
||||
oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
and oauth_client_metadata.token_endpoint_auth_method
|
||||
not in oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
):
|
||||
oauth_client_metadata.token_endpoint_auth_method = (
|
||||
oauth_server_metadata.token_endpoint_auth_methods_supported[0]
|
||||
)
|
||||
|
||||
registration_url = None
|
||||
if oauth_server_metadata and oauth_server_metadata.registration_endpoint:
|
||||
if oauth_server_metadata.registration_endpoint:
|
||||
registration_url = str(oauth_server_metadata.registration_endpoint)
|
||||
else:
|
||||
_, base_url = get_parsed_and_base_url(oauth_server_url)
|
||||
|
|
@ -661,9 +654,6 @@ async def get_oauth_client_info_with_static_credentials(
|
|||
but skips dynamic client registration entirely.
|
||||
"""
|
||||
try:
|
||||
oauth_server_metadata = None
|
||||
oauth_server_metadata_url = None
|
||||
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
redirect_uri = f'{redirect_base_url}/oauth/clients/{client_id}/callback'
|
||||
|
|
@ -671,18 +661,9 @@ async def get_oauth_client_info_with_static_credentials(
|
|||
# Discover server metadata (authorization endpoint, token endpoint, scopes, etc.)
|
||||
resource_metadata = await get_protected_resource_metadata(oauth_server_url)
|
||||
resource = resource_metadata.resource
|
||||
discovery_urls = resource_metadata.get_discovery_urls(oauth_server_url)
|
||||
for url in discovery_urls:
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as resp:
|
||||
if resp.status == 200:
|
||||
try:
|
||||
oauth_server_metadata = OAuthMetadata.model_validate(await resp.json())
|
||||
oauth_server_metadata_url = url
|
||||
break
|
||||
except Exception as e:
|
||||
log.error(f'Error parsing OAuth metadata from {url}: {e}')
|
||||
continue
|
||||
oauth_server_metadata, oauth_server_metadata_url = await _get_oauth_server_metadata(
|
||||
oauth_server_url, resource_metadata
|
||||
)
|
||||
|
||||
# Use scopes from the Protected Resource Metadata (RFC 9728) if available.
|
||||
# Unlike the Authorization Server's scopes_supported (which is a full catalog
|
||||
|
|
@ -695,8 +676,7 @@ async def get_oauth_client_info_with_static_credentials(
|
|||
# Determine token_endpoint_auth_method
|
||||
token_endpoint_auth_method = 'client_secret_post'
|
||||
if (
|
||||
oauth_server_metadata
|
||||
and oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
and token_endpoint_auth_method not in oauth_server_metadata.token_endpoint_auth_methods_supported
|
||||
):
|
||||
token_endpoint_auth_method = oauth_server_metadata.token_endpoint_auth_methods_supported[0]
|
||||
|
|
@ -1149,14 +1129,8 @@ class OAuthClientManager:
|
|||
log.error(f'No OAuth client found for provider {client_id}')
|
||||
return None
|
||||
|
||||
token_endpoint = None
|
||||
async with aiohttp.ClientSession(trust_env=True) as session_http:
|
||||
async with session_http.get(await self.get_server_metadata_url(client_id)) as r:
|
||||
if r.status == 200:
|
||||
openid_data = await r.json()
|
||||
token_endpoint = openid_data.get('token_endpoint')
|
||||
else:
|
||||
log.error(f'Failed to fetch OpenID configuration for client_id {client_id}')
|
||||
metadata = await client.load_server_metadata()
|
||||
token_endpoint = metadata.get('token_endpoint') or client.access_token_url
|
||||
if not token_endpoint:
|
||||
log.error(f'No token endpoint found for client_id {client_id}')
|
||||
return None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue