Coverage for open_webui/utils/oauth.py: 15%
1208 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import asyncio
2import base64
3import fnmatch
4import hashlib
5import logging
6import re
7import sys
8import urllib
9import uuid
10from dataclasses import dataclass, field
11from datetime import datetime, timedelta
12from functools import partialmethod
13from types import SimpleNamespace
14from typing import Literal, Optional
16import aiohttp
17import jwt
18from authlib.integrations.starlette_client import OAuth
19from authlib.oauth2.rfc6749.errors import OAuth2Error
20from authlib.oidc.core import UserInfo
21from cryptography.fernet import Fernet, InvalidToken
22from fastapi import (
23 HTTPException,
24 status,
25)
26from joserfc.errors import BadSignatureError
27from joserfc.jws import JWSRegistry
28from mcp.shared.auth import (
29 OAuthClientMetadata as MCPOAuthClientMetadata,
30)
31from mcp.shared.auth import (
32 OAuthMetadata,
33)
34from open_webui.config import (
35 DEFAULT_USER_ROLE,
36 ENABLE_OAUTH,
37 ENABLE_OAUTH_GROUP_CREATION,
38 ENABLE_OAUTH_GROUP_MANAGEMENT,
39 ENABLE_OAUTH_ROLE_MANAGEMENT,
40 ENABLE_OAUTH_SIGNUP,
41 JWT_EXPIRES_IN,
42 OAUTH_ACCESS_TOKEN_REQUEST_INCLUDE_CLIENT_ID,
43 OAUTH_ADMIN_ROLES,
44 OAUTH_ALLOWED_DOMAINS,
45 OAUTH_ALLOWED_ROLES,
46 OAUTH_AUDIENCE,
47 OAUTH_AUTHORIZE_PARAMS,
48 OAUTH_BLOCKED_GROUPS,
49 OAUTH_CLIENT_TIMEOUT,
50 OAUTH_EMAIL_CLAIM,
51 OAUTH_GROUP_DEFAULT_SHARE,
52 OAUTH_GROUPS_CLAIM,
53 OAUTH_GROUPS_SEPARATOR,
54 OAUTH_MERGE_ACCOUNTS_BY_EMAIL,
55 OAUTH_PICTURE_CLAIM,
56 OAUTH_PROVIDERS,
57 OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE,
58 OAUTH_ROLES_CLAIM,
59 OAUTH_ROLES_SEPARATOR,
60 OAUTH_SUB_CLAIM,
61 OAUTH_UPDATE_EMAIL_ON_LOGIN,
62 OAUTH_UPDATE_NAME_ON_LOGIN,
63 OAUTH_UPDATE_PICTURE_ON_LOGIN,
64 OAUTH_USERNAME_CLAIM,
65 WEBHOOK_URL,
66)
67from open_webui.constants import ERROR_MESSAGES
68from open_webui.env import (
69 AIOHTTP_CLIENT_ALLOW_REDIRECTS,
70 AIOHTTP_CLIENT_SESSION_SSL,
71 ENABLE_OAUTH_EMAIL_FALLBACK,
72 ENABLE_OAUTH_ID_TOKEN_COOKIE,
73 OAUTH_CLIENT_INFO_ENCRYPTION_KEY,
74 OAUTH_MAX_SESSIONS_PER_USER,
75 WEBUI_AUTH_COOKIE_SAME_SITE,
76 WEBUI_AUTH_COOKIE_SECURE,
77)
78from open_webui.events import EVENTS, publish_event
79from open_webui.models.auths import Auths
80from open_webui.models.config import Config
81from open_webui.models.groups import GroupForm, GroupModel, Groups, GroupUpdateForm
82from open_webui.models.oauth_sessions import OAuthSessions
83from open_webui.models.users import Users
84from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url
85from open_webui.utils.auth import (
86 create_token,
87 get_password_hash,
88 get_optional_verified_user_from_request,
89 get_verified_user_by_id,
90 revoke_user_tokens,
91)
92from open_webui.utils.groups import apply_default_group_assignment
93from open_webui.utils.misc import parse_duration
94from open_webui.utils.validate import validate_image_url
95from starlette.responses import RedirectResponse
97# Some IdPs put private params in ID token JOSE headers (CAS: client_id, CyberArk: app_id).
98# Authlib exposes no way to pass a registry, so relax it globally; crit, alg and signature checks still apply.
99JWSRegistry.__init__ = partialmethod(JWSRegistry.__init__, strict_check_header=False)
102class OAuthClientMetadata(MCPOAuthClientMetadata):
103 token_endpoint_auth_method: Literal['none', 'client_secret_basic', 'client_secret_post'] = 'client_secret_post'
104 pass
107OAuthResourceParameterMode = Literal['auto', 'include', 'omit']
110class OAuthClientInformationFull(OAuthClientMetadata):
111 issuer: Optional[str] = None # URL of the OAuth server that issued this client
112 resource: Optional[str] = None # RFC 8707 resource indicator for JWT audience
113 oauth_resource_parameter: OAuthResourceParameterMode = 'auto'
115 client_id: str
116 client_secret: str | None = None
117 client_id_issued_at: int | None = None
118 client_secret_expires_at: int | None = None
120 server_metadata: Optional[OAuthMetadata] = None # Fetched from the OAuth server
123from open_webui.env import GLOBAL_LOG_LEVEL
124from open_webui.utils.json_codec import JSONCodec
126logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
127log = logging.getLogger(__name__)
129OAUTH_RESOURCE_PARAMETER_MODES = {'auto', 'include', 'omit'}
131OAUTH_RUNTIME_CONFIG = {
132 'DEFAULT_USER_ROLE': ('ui.default_user_role', DEFAULT_USER_ROLE),
133 'ENABLE_OAUTH': ('oauth.enable', ENABLE_OAUTH),
134 'ENABLE_OAUTH_SIGNUP': ('oauth.enable_signup', ENABLE_OAUTH_SIGNUP),
135 'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE': (
136 'oauth.refresh_token.include_scope',
137 OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE,
138 ),
139 'OAUTH_MERGE_ACCOUNTS_BY_EMAIL': (
140 'oauth.merge_accounts_by_email',
141 OAUTH_MERGE_ACCOUNTS_BY_EMAIL,
142 ),
143 'ENABLE_OAUTH_ROLE_MANAGEMENT': (
144 'oauth.enable_role_mapping',
145 ENABLE_OAUTH_ROLE_MANAGEMENT,
146 ),
147 'ENABLE_OAUTH_GROUP_MANAGEMENT': (
148 'oauth.enable_group_mapping',
149 ENABLE_OAUTH_GROUP_MANAGEMENT,
150 ),
151 'ENABLE_OAUTH_GROUP_CREATION': (
152 'oauth.enable_group_creation',
153 ENABLE_OAUTH_GROUP_CREATION,
154 ),
155 'OAUTH_GROUP_DEFAULT_SHARE': (
156 'oauth.group_default_share',
157 OAUTH_GROUP_DEFAULT_SHARE,
158 ),
159 'OAUTH_BLOCKED_GROUPS': ('oauth.blocked_groups', OAUTH_BLOCKED_GROUPS),
160 'OAUTH_ROLES_CLAIM': ('oauth.roles_claim', OAUTH_ROLES_CLAIM),
161 'OAUTH_SUB_CLAIM': ('oauth.sub_claim', OAUTH_SUB_CLAIM),
162 'OAUTH_GROUPS_CLAIM': ('oauth.group_claim', OAUTH_GROUPS_CLAIM),
163 'OAUTH_EMAIL_CLAIM': ('oauth.email_claim', OAUTH_EMAIL_CLAIM),
164 'OAUTH_PICTURE_CLAIM': ('oauth.picture_claim', OAUTH_PICTURE_CLAIM),
165 'OAUTH_USERNAME_CLAIM': ('oauth.username_claim', OAUTH_USERNAME_CLAIM),
166 'OAUTH_ALLOWED_ROLES': ('oauth.allowed_roles', OAUTH_ALLOWED_ROLES),
167 'OAUTH_ADMIN_ROLES': ('oauth.admin_roles', OAUTH_ADMIN_ROLES),
168 'OAUTH_ALLOWED_DOMAINS': ('oauth.allowed_domains', OAUTH_ALLOWED_DOMAINS),
169 'WEBHOOK_URL': ('webhook_url', WEBHOOK_URL),
170 'JWT_EXPIRES_IN': ('auth.jwt_expiry', JWT_EXPIRES_IN),
171 'OAUTH_UPDATE_PICTURE_ON_LOGIN': (
172 'oauth.update_picture_on_login',
173 OAUTH_UPDATE_PICTURE_ON_LOGIN,
174 ),
175 'OAUTH_UPDATE_NAME_ON_LOGIN': (
176 'oauth.update_name_on_login',
177 OAUTH_UPDATE_NAME_ON_LOGIN,
178 ),
179 'OAUTH_UPDATE_EMAIL_ON_LOGIN': (
180 'oauth.update_email_on_login',
181 OAUTH_UPDATE_EMAIL_ON_LOGIN,
182 ),
183 'OAUTH_AUDIENCE': ('oauth.audience', OAUTH_AUDIENCE),
184}
187def _default_value(value):
188 return getattr(value, 'value', value)
191def _get_roles_claim(claims: dict, claim: str) -> list | str | int | None:
192 """Read nested or flat claims, preserving explicit empty values and zero."""
193 value = claims
194 for key in claim.split('.'):
195 value = value.get(key) if isinstance(value, dict) else None
196 if not isinstance(value, (list, str, int)):
197 value = claims.get(claim)
198 return value if isinstance(value, (list, str, int)) else None
201async def get_oauth_runtime_config() -> SimpleNamespace:
202 keys = [key for key, _default in OAUTH_RUNTIME_CONFIG.values()]
203 stored = await Config.get_many(*keys)
204 values = {name: stored.get(key, _default_value(default)) for name, (key, default) in OAUTH_RUNTIME_CONFIG.items()}
205 return SimpleNamespace(**values)
208# Conservative default when the provider omits both expires_in and expires_at.
209# Matches the value recommended by Authlib's compliance_fix documentation.
210DEFAULT_TOKEN_EXPIRY_SECONDS = 3600
211NON_EXPIRING_TOKEN_EXPIRES_AT = 253402300799 # 9999-12-31 23:59:59 UTC
214def _normalize_token_expiry(token: dict) -> dict:
215 """Ensure a token dict always has a numeric access-token ``expires_at``.
217 Resolution order:
218 1. If *expires_at* is already present and non-None, trust it.
219 2. Else if *expires_in* is present and non-None, compute *expires_at*.
220 3. Else if a *refresh_token* is present, fall back to
221 ``DEFAULT_TOKEN_EXPIRY_SECONDS`` and log a warning so operators can
222 identify providers that omit expiration.
223 4. Otherwise treat the token as non-expiring; there is no refresh path to
224 recover from a fabricated short expiry.
226 Also stamps *issued_at* for auditing.
227 """
228 token['issued_at'] = datetime.now().timestamp()
230 if token.get('expires_at') is not None:
231 expires_at = int(token['expires_at'])
232 elif token.get('expires_in') is not None:
233 expires_at = int(datetime.now().timestamp() + token['expires_in'])
234 elif token.get('refresh_token'):
235 log.warning(
236 "OAuth token response missing both 'expires_in' and 'expires_at'; "
237 f'defaulting to {DEFAULT_TOKEN_EXPIRY_SECONDS}s from now'
238 )
239 expires_at = int(datetime.now().timestamp() + DEFAULT_TOKEN_EXPIRY_SECONDS)
240 else:
241 log.info(
242 "OAuth token response missing 'expires_in', 'expires_at' and 'refresh_token'; treating token as non-expiring"
243 )
244 expires_at = NON_EXPIRING_TOKEN_EXPIRES_AT
246 token['expires_at'] = expires_at
247 return token
250FERNET = None
252if len(OAUTH_CLIENT_INFO_ENCRYPTION_KEY) != 44: 252 ↛ 256line 252 didn't jump to line 256 because the condition on line 252 was always true
253 key_bytes = hashlib.sha256(OAUTH_CLIENT_INFO_ENCRYPTION_KEY.encode()).digest()
254 OAUTH_CLIENT_INFO_ENCRYPTION_KEY = base64.urlsafe_b64encode(key_bytes)
255else:
256 OAUTH_CLIENT_INFO_ENCRYPTION_KEY = OAUTH_CLIENT_INFO_ENCRYPTION_KEY.encode()
258try:
259 FERNET = Fernet(OAUTH_CLIENT_INFO_ENCRYPTION_KEY)
260except Exception as e:
261 log.error(f'Error initializing Fernet with provided key: {e}')
262 raise
265def encrypt_data(data) -> str:
266 """Encrypt data for storage"""
267 try:
268 data_json = JSONCodec.dumps(data)
269 encrypted = FERNET.encrypt(data_json.encode()).decode()
270 return encrypted
271 except Exception as e:
272 log.error(f'Error encrypting data: {e}')
273 raise
276def decrypt_data(data: str):
277 """Decrypt data from storage"""
278 decrypted = FERNET.decrypt(data.encode()).decode()
279 return JSONCodec.loads(decrypted)
282def _build_oauth_callback_error_message(e: Exception) -> str:
283 """
284 Produce a user-facing callback error string with actionable context.
285 Keeps the message short and strips newlines for safe redirect usage.
286 """
287 if isinstance(e, OAuth2Error):
288 parts = [p for p in [e.error, e.description] if p]
289 detail = ' - '.join(parts)
290 elif isinstance(e, HTTPException):
291 detail = e.detail if isinstance(e.detail, str) else str(e.detail)
292 elif isinstance(e, aiohttp.ClientResponseError):
293 detail = f'Upstream provider returned {e.status}: {e.message}'
294 elif isinstance(e, aiohttp.ClientError):
295 detail = str(e)
296 elif isinstance(e, KeyError):
297 missing = str(e).strip("'")
298 if missing.lower() == 'state':
299 detail = 'Missing state parameter in callback (session may have expired)'
300 else:
301 detail = f"Missing expected key '{missing}' in OAuth response"
302 else:
303 detail = str(e)
305 detail = detail.replace('\n', ' ').strip()
306 if not detail:
307 detail = e.__class__.__name__
309 message = f'OAuth callback failed: {detail}'
310 return message[:197] + '...' if len(message) > 200 else message
313def is_in_blocked_groups(group_name: str, groups: list) -> bool:
314 """
315 Check if a group name matches any blocked pattern.
316 Supports exact matches, shell-style wildcards (*, ?), and regex patterns.
318 Args:
319 group_name: The group name to check
320 groups: List of patterns to match against
322 Returns:
323 True if the group is blocked, False otherwise
324 """
325 if not groups:
326 return False
328 for group_pattern in groups:
329 if not group_pattern: # Skip empty patterns
330 continue
332 # Exact match
333 if group_name == group_pattern:
334 return True
336 # Try as regex pattern first if it contains regex-specific characters
337 if any(char in group_pattern for char in ['^', '$', '[', ']', '(', ')', '{', '}', '+', '\\', '|']):
338 try:
339 # Use the original pattern as-is for regex matching
340 if re.search(group_pattern, group_name):
341 return True
342 except re.error:
343 # If regex is invalid, fall through to wildcard check
344 pass
346 # Shell-style wildcard match (supports * and ?)
347 if '*' in group_pattern or '?' in group_pattern:
348 if fnmatch.fnmatch(group_name, group_pattern):
349 return True
351 return False
354def _parse_blocked_groups(value) -> list[str]:
355 """Accept JSON arrays, persisted lists, and comma-separated admin input."""
356 if isinstance(value, str):
357 try:
358 parsed = JSONCodec.loads(value)
359 except JSONCodec.JSONDecodeError:
360 parsed = None
361 value = parsed if isinstance(parsed, list) else [group.strip() for group in value.split(',')]
362 if not isinstance(value, list):
363 return []
364 return [group for group in value if isinstance(group, str) and group]
367def get_parsed_and_base_url(server_url) -> tuple[urllib.parse.ParseResult, str]:
368 parsed = urllib.parse.urlparse(server_url)
369 base_url = f'{parsed.scheme}://{parsed.netloc}'
370 return parsed, base_url
373@dataclass
374class ProtectedResourceMetadata:
375 """RFC 9728 Protected Resource Metadata fields relevant to OAuth flows."""
377 resource: str | None = None
378 authorization_servers: list[str] = field(default_factory=list)
379 scopes_supported: list[str] = field(default_factory=list)
381 def get_discovery_urls(self, server_url: str) -> list[str]:
382 """Build all candidate OAuth discovery URLs from this metadata and the server URL."""
383 urls = []
384 for auth_server in self.authorization_servers: 384 ↛ 385line 384 didn't jump to line 385 because the loop on line 384 never started
385 urls.extend(_build_well_known_urls(auth_server.rstrip('/')))
386 urls.extend(_build_well_known_urls(server_url))
387 return urls
390async def get_protected_resource_metadata(server_url: str) -> ProtectedResourceMetadata:
391 """
392 Fetch RFC 9728 Protected Resource Metadata from an MCP server.
394 https://modelcontextprotocol.io/specification/2025-03-26/basic/authorization
396 Returns:
397 ProtectedResourceMetadata with the resource indicator (RFC 8707)
398 and authorization server URLs discovered from the metadata document.
399 """
400 authorization_servers = []
401 resource = None
402 scopes = []
403 try:
404 async with aiohttp.ClientSession(trust_env=True) as session:
405 async with session.post(
406 server_url,
407 json={'jsonrpc': '2.0', 'method': 'initialize', 'params': {}, 'id': 1},
408 headers={'Content-Type': 'application/json'},
409 ssl=AIOHTTP_CLIENT_SESSION_SSL,
410 ) as response:
411 # Discover Protected Resource Metadata regardless of HTTP status.
412 # A 401 carries a WWW-Authenticate header pointing at the PRM, but
413 # some MCP servers (e.g. Google's gmail/drive/calendar remote MCPs)
414 # answer 200 to an anonymous `initialize`, so we must still fall
415 # back to the RFC 9728 well-known URIs when there is no 401/header.
416 resource_metadata_urls = []
417 match = re.search(
418 r'resource_metadata=(?:"([^"]+)"|([^\s,]+))',
419 response.headers.get('WWW-Authenticate', ''),
420 )
421 if match:
422 resource_metadata_urls = [match.group(1) or match.group(2)]
423 log.debug('Found resource_metadata URL: %s', resource_metadata_urls[0])
424 else:
425 # Fall back to well-known resource metadata URIs (RFC 9728 §4.2)
426 parsed, base_url = get_parsed_and_base_url(server_url)
427 if parsed.path and parsed.path != '/':
428 path = parsed.path.rstrip('/')
429 resource_metadata_urls.append(
430 urllib.parse.urljoin(base_url, f'/.well-known/oauth-protected-resource{path}')
431 )
432 resource_metadata_urls.append(
433 urllib.parse.urljoin(base_url, '/.well-known/oauth-protected-resource')
434 )
435 log.debug('No resource_metadata in header, trying well-known URIs: %s', resource_metadata_urls)
437 # Fetch Protected Resource metadata from candidate URLs
438 for resource_metadata_url in resource_metadata_urls:
439 try:
440 async with session.get(
441 resource_metadata_url, ssl=AIOHTTP_CLIENT_SESSION_SSL
442 ) as resource_response:
443 if resource_response.status == 200:
444 resource_metadata = await resource_response.json()
446 resource = resource_metadata.get('resource') or None
447 if resource:
448 log.debug('Discovered resource indicator: %s', resource)
450 servers = resource_metadata.get('authorization_servers', [])
451 scopes = resource_metadata.get('scopes_supported', [])
452 if scopes:
453 log.debug('Discovered resource scopes: %s', scopes)
455 if servers:
456 authorization_servers = servers
457 log.debug('Discovered authorization servers: %s', servers)
458 break
459 except Exception as e:
460 log.debug('Failed to fetch resource metadata from %s: %s', resource_metadata_url, e)
461 continue
462 except Exception as e:
463 log.debug('MCP Protected Resource discovery failed: %s', e)
465 return ProtectedResourceMetadata(
466 resource=resource, authorization_servers=authorization_servers, scopes_supported=scopes
467 )
470def _build_well_known_urls(server_url: str) -> list[str]:
471 """Build RFC 8414 / OIDC Discovery well-known URLs for a server URL."""
472 parsed, base_url = get_parsed_and_base_url(server_url)
473 urls = []
475 if parsed.path and parsed.path != '/':
476 path = parsed.path.rstrip('/')
477 urls.extend(
478 [
479 urllib.parse.urljoin(base_url, f'/.well-known/oauth-authorization-server{path}'),
480 urllib.parse.urljoin(base_url, f'/.well-known/openid-configuration{path}'),
481 urllib.parse.urljoin(base_url, f'{path}/.well-known/openid-configuration'),
482 ]
483 )
485 urls.extend(
486 [
487 urllib.parse.urljoin(base_url, '/.well-known/oauth-authorization-server'),
488 urllib.parse.urljoin(base_url, '/.well-known/openid-configuration'),
489 ]
490 )
492 return urls
495async def get_discovery_urls(server_url) -> list[str]:
496 """Convenience: get all OAuth discovery URLs for a server URL."""
497 metadata = await get_protected_resource_metadata(server_url)
498 return metadata.get_discovery_urls(server_url)
501# TODO: Some OAuth providers require Initial Access Tokens (IATs) for dynamic client registration.
502# This is not currently supported.
503async def get_oauth_client_info_with_dynamic_client_registration(
504 request,
505 client_id: str,
506 oauth_server_url: str,
507 oauth_server_key: Optional[str] = None,
508 oauth_scope: str | None = None,
509) -> OAuthClientInformationFull:
510 try:
511 oauth_server_metadata = None
512 oauth_server_metadata_url = None
514 webui_url = await Config.get('webui.url')
515 redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
517 oauth_client_metadata = OAuthClientMetadata(
518 # LICENSE covers this Open WebUI OAuth client identifier.
519 # Do not alter, remove, obscure, or replace it except as LICENSE permits:
520 # https://docs.openwebui.com/license.
521 client_name='Open WebUI',
522 redirect_uris=[f'{redirect_base_url}/oauth/clients/{client_id}/callback'],
523 grant_types=['authorization_code', 'refresh_token'],
524 response_types=['code'],
525 )
527 # Attempt to fetch OAuth server metadata to get registration endpoint & scopes
528 resource_metadata = await get_protected_resource_metadata(oauth_server_url)
529 resource = resource_metadata.resource
531 # Prefer the resource-specific scopes from the Protected Resource Metadata
532 # (RFC 9728) over the AS's full scopes_supported catalog, for least
533 # privilege. Mirrors the static-credentials flow (#24690).
534 scope_override = ' '.join(oauth_scope.replace(',', ' ').split()) if oauth_scope else None
535 if scope_override:
536 oauth_client_metadata.scope = scope_override
537 elif resource_metadata.scopes_supported: 537 ↛ 538line 537 didn't jump to line 538 because the condition on line 537 was never true
538 oauth_client_metadata.scope = ' '.join(resource_metadata.scopes_supported)
540 discovery_urls = resource_metadata.get_discovery_urls(oauth_server_url)
541 for url in discovery_urls: 541 ↛ 576line 541 didn't jump to line 576 because the loop on line 541 didn't complete
542 async with aiohttp.ClientSession(trust_env=True) as session:
543 async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as oauth_server_metadata_response:
544 if oauth_server_metadata_response.status == 200:
545 try:
546 oauth_server_metadata = OAuthMetadata.model_validate(
547 await oauth_server_metadata_response.json()
548 )
549 oauth_server_metadata_url = url
550 if (
551 oauth_client_metadata.scope is None
552 and oauth_server_metadata.scopes_supported is not None
553 ):
554 oauth_client_metadata.scope = ' '.join(oauth_server_metadata.scopes_supported)
556 if (
557 oauth_server_metadata.token_endpoint_auth_methods_supported
558 and oauth_client_metadata.token_endpoint_auth_method
559 not in oauth_server_metadata.token_endpoint_auth_methods_supported
560 ):
561 # Pick the first supported method from the server
562 oauth_client_metadata.token_endpoint_auth_method = (
563 oauth_server_metadata.token_endpoint_auth_methods_supported[0]
564 )
566 break
567 except Exception as e:
568 log.error(f'Error parsing OAuth metadata from {url}: {e}')
569 continue
571 # Fail fast if authorization server metadata discovery did not resolve an
572 # authorization endpoint. Otherwise registration can still "succeed" (via
573 # the /register fallback below) while issuer/server_metadata stay unset,
574 # which later crashes at authorize time with authlib's
575 # RuntimeError: Missing "authorize_url" value. (#26647)
576 if oauth_server_metadata is None or not oauth_server_metadata.authorization_endpoint:
577 log.error(f'OAuth authorization server metadata discovery failed for {oauth_server_url}')
578 raise Exception(
579 'Could not discover the OAuth authorization server metadata '
580 f'(authorization_endpoint) for {oauth_server_url}. The MCP server must '
581 'expose RFC 8414 / RFC 9728 discovery documents so Open WebUI can '
582 'resolve where to send users to authorize.'
583 )
585 registration_url = None
586 if oauth_server_metadata and oauth_server_metadata.registration_endpoint:
587 registration_url = str(oauth_server_metadata.registration_endpoint)
588 else:
589 _, base_url = get_parsed_and_base_url(oauth_server_url)
590 registration_url = urllib.parse.urljoin(base_url, '/register')
592 registration_data = oauth_client_metadata.model_dump(
593 exclude_none=True,
594 mode='json',
595 by_alias=True,
596 )
598 # Perform dynamic client registration and return client info
599 async with aiohttp.ClientSession(trust_env=True) as session:
600 async with session.post(
601 registration_url, json=registration_data, ssl=AIOHTTP_CLIENT_SESSION_SSL
602 ) as oauth_client_registration_response:
603 try:
604 registration_response_json = await oauth_client_registration_response.json()
606 # The mcp package requires optional unset values to be None. If an empty string is passed, it gets validated and fails.
607 # This replaces all empty strings with None.
608 registration_response_json = {
609 k: (None if v == '' else v) for k, v in registration_response_json.items()
610 }
611 oauth_client_info = OAuthClientInformationFull.model_validate(
612 {
613 **registration_response_json,
614 'issuer': oauth_server_metadata_url,
615 'server_metadata': oauth_server_metadata,
616 'resource': resource,
617 }
618 )
619 log.info(
620 'Dynamic client registration successful at %s, client_id: %s',
621 registration_url,
622 oauth_client_info.client_id,
623 )
624 return oauth_client_info
625 except Exception as e:
626 error_text = None
627 try:
628 error_text = await oauth_client_registration_response.text()
629 log.error(
630 f'Dynamic client registration failed at {registration_url}: {oauth_client_registration_response.status} - {error_text}'
631 )
632 except Exception as e:
633 pass
635 log.error(f'Error parsing client registration response: {e}')
636 raise Exception(
637 f'Dynamic client registration failed: {error_text}'
638 if error_text
639 else 'Error parsing client registration response'
640 )
641 raise Exception('Dynamic client registration failed')
642 except Exception as e:
643 log.error(f'Exception during dynamic client registration: {e}')
644 raise e
647async def get_oauth_client_info_with_static_credentials(
648 request,
649 client_id: str,
650 oauth_server_url: str,
651 oauth_client_id: str,
652 oauth_client_secret: str,
653 oauth_scope: str | None = None,
654) -> OAuthClientInformationFull:
655 """
656 Build an OAuthClientInformationFull from user-provided static credentials.
657 Performs server metadata discovery to resolve authorization/token endpoints,
658 but skips dynamic client registration entirely.
659 """
660 try:
661 oauth_server_metadata = None
662 oauth_server_metadata_url = None
664 webui_url = await Config.get('webui.url')
665 redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
666 redirect_uri = f'{redirect_base_url}/oauth/clients/{client_id}/callback'
668 # Discover server metadata (authorization endpoint, token endpoint, scopes, etc.)
669 resource_metadata = await get_protected_resource_metadata(oauth_server_url)
670 resource = resource_metadata.resource
671 discovery_urls = resource_metadata.get_discovery_urls(oauth_server_url)
672 for url in discovery_urls: 672 ↛ 688line 672 didn't jump to line 688 because the loop on line 672 didn't complete
673 async with aiohttp.ClientSession(trust_env=True) as session:
674 async with session.get(url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as resp:
675 if resp.status == 200:
676 try:
677 oauth_server_metadata = OAuthMetadata.model_validate(await resp.json())
678 oauth_server_metadata_url = url
679 break
680 except Exception as e:
681 log.error(f'Error parsing OAuth metadata from {url}: {e}')
682 continue
684 # Use scopes from the Protected Resource Metadata (RFC 9728) if available.
685 # Unlike the Authorization Server's scopes_supported (which is a full catalog
686 # of every scope the server can grant), the PRM scopes_supported represents
687 # what this specific resource requires — making it safe to request them all.
688 scope = (' '.join(oauth_scope.replace(',', ' ').split()) if oauth_scope else None) or (
689 ' '.join(resource_metadata.scopes_supported) if resource_metadata.scopes_supported else None
690 )
692 # Determine token_endpoint_auth_method
693 token_endpoint_auth_method = 'client_secret_post'
694 if (
695 oauth_server_metadata
696 and oauth_server_metadata.token_endpoint_auth_methods_supported
697 and token_endpoint_auth_method not in oauth_server_metadata.token_endpoint_auth_methods_supported
698 ):
699 token_endpoint_auth_method = oauth_server_metadata.token_endpoint_auth_methods_supported[0]
701 oauth_client_info = OAuthClientInformationFull(
702 client_id=oauth_client_id,
703 client_secret=oauth_client_secret,
704 redirect_uris=[redirect_uri],
705 grant_types=['authorization_code', 'refresh_token'],
706 response_types=['code'],
707 scope=scope,
708 token_endpoint_auth_method=token_endpoint_auth_method,
709 issuer=oauth_server_metadata_url,
710 server_metadata=oauth_server_metadata,
711 resource=resource,
712 )
714 log.info(
715 'Static OAuth client info built for %s using metadata from %s', oauth_client_id, oauth_server_metadata_url
716 )
717 return oauth_client_info
718 except Exception as e:
719 log.error(f'Exception building static OAuth client info: {e}')
720 raise e
723def resolve_oauth_client_info(connection: dict) -> dict:
724 """
725 Decrypt OAuth client info from a tool server connection config.
727 For oauth_2.1_static, overlays admin-provided credentials from
728 info.oauth_client_id and info.oauth_client_secret onto the blob.
729 """
730 info = connection.get('info') or {}
731 data = decrypt_data(info.get('oauth_client_info', ''))
733 if connection.get('auth_type') == 'oauth_2.1_static':
734 if info.get('oauth_client_id') and info.get('oauth_client_secret'):
735 data['client_id'] = info['oauth_client_id']
736 data['client_secret'] = info['oauth_client_secret']
738 return data
741def normalize_oauth_resource_parameter(value: str | None) -> OAuthResourceParameterMode:
742 if value in OAUTH_RESOURCE_PARAMETER_MODES:
743 return value
744 return 'auto'
747def get_connection_oauth_resource_parameter(connection: dict) -> OAuthResourceParameterMode:
748 info = connection.get('info') or {}
749 config = connection.get('config') or {}
750 return normalize_oauth_resource_parameter(
751 info.get('oauth_resource_parameter') or config.get('oauth_resource_parameter')
752 )
755def apply_connection_oauth_options(connection: dict, oauth_client_info: dict) -> dict:
756 info = connection.get('info') or {}
757 config = connection.get('config') or {}
758 oauth_scope = info.get('oauth_scope') or config.get('oauth_scope')
759 oauth_scope = ' '.join(oauth_scope.replace(',', ' ').split()) if oauth_scope else None
761 options = {
762 **oauth_client_info,
763 'oauth_resource_parameter': get_connection_oauth_resource_parameter(connection),
764 }
765 if oauth_scope:
766 options['scope'] = oauth_scope
767 return options
770def scope_has_resource_indicator(scope: str | None) -> bool:
771 if not scope:
772 return False
773 return any(scope_value.startswith(('https://', 'http://', 'api://')) for scope_value in scope.split())
776def should_send_oauth_resource(client_info: OAuthClientInformationFull | None) -> bool:
777 if not client_info or not client_info.resource:
778 return False
780 mode = normalize_oauth_resource_parameter(client_info.oauth_resource_parameter)
781 if mode == 'omit':
782 return False
783 if mode == 'include':
784 return True
786 return not scope_has_resource_indicator(client_info.scope)
789def build_oauth_request_params(client_info: OAuthClientInformationFull | None) -> dict:
790 if not client_info:
791 return {}
793 params = {}
794 if client_info.scope:
795 params['scope'] = client_info.scope
796 if should_send_oauth_resource(client_info):
797 params['resource'] = client_info.resource
798 return params
801async def recover_static_oauth_client_metadata(connection: dict, oauth_client_info: dict) -> dict:
802 if connection.get('auth_type') != 'oauth_2.1_static':
803 return oauth_client_info
805 if oauth_client_info.get('scope') and oauth_client_info.get('resource'):
806 return oauth_client_info
808 server_url = connection.get('url')
809 if not server_url:
810 return oauth_client_info
812 try:
813 resource_metadata = await get_protected_resource_metadata(server_url)
814 except Exception as e:
815 log.debug('Unable to recover static OAuth metadata for %s: %s', server_url, e)
816 return oauth_client_info
818 recovered = {**oauth_client_info}
819 if not recovered.get('scope') and resource_metadata.scopes_supported:
820 recovered['scope'] = ' '.join(resource_metadata.scopes_supported)
821 log.info('Recovered static OAuth scopes for %s from protected resource metadata', server_url)
823 if not recovered.get('resource') and resource_metadata.resource:
824 recovered['resource'] = resource_metadata.resource
826 return recovered
829class OAuthClientManager:
830 def __init__(self, app):
831 self.oauth = OAuth()
832 self.app = app
833 self.clients = {}
835 def add_client(self, client_id, oauth_client_info: OAuthClientInformationFull):
836 kwargs = {
837 'name': client_id,
838 'client_id': oauth_client_info.client_id,
839 'client_secret': oauth_client_info.client_secret,
840 'client_kwargs': {
841 'follow_redirects': True,
842 **({'timeout': int(OAUTH_CLIENT_TIMEOUT)} if OAUTH_CLIENT_TIMEOUT else {}),
843 **({'scope': oauth_client_info.scope} if oauth_client_info.scope else {}),
844 **(
845 {'token_endpoint_auth_method': oauth_client_info.token_endpoint_auth_method}
846 if oauth_client_info.token_endpoint_auth_method
847 else {}
848 ),
849 },
850 'server_metadata_url': (oauth_client_info.issuer if oauth_client_info.issuer else None),
851 }
853 # Defense-in-depth: when the server metadata is already known, pass the
854 # authorization/token endpoints explicitly so authlib does not rely solely
855 # on refetching server_metadata_url (which may be missing/unreachable) to
856 # resolve them. Prevents RuntimeError: Missing "authorize_url". (#26647)
857 server_metadata = oauth_client_info.server_metadata
858 if server_metadata is not None:
859 if getattr(server_metadata, 'authorization_endpoint', None):
860 kwargs['authorize_url'] = str(server_metadata.authorization_endpoint)
861 if getattr(server_metadata, 'token_endpoint', None):
862 kwargs['access_token_url'] = str(server_metadata.token_endpoint)
864 # Default to S256 for OAuth 2.1 (PKCE is mandatory per RFC 9700)
865 kwargs['code_challenge_method'] = 'S256'
867 # Only remove PKCE if metadata explicitly excludes S256
868 if (
869 oauth_client_info.server_metadata
870 and oauth_client_info.server_metadata.code_challenge_methods_supported
871 and isinstance(
872 oauth_client_info.server_metadata.code_challenge_methods_supported,
873 list,
874 )
875 and 'S256' not in oauth_client_info.server_metadata.code_challenge_methods_supported
876 ):
877 del kwargs['code_challenge_method']
879 self.clients[client_id] = {
880 'client': self.oauth.register(**kwargs),
881 'client_info': oauth_client_info,
882 }
883 return self.clients[client_id]
885 async def ensure_client_from_config(self, client_id):
886 """
887 Lazy-load an OAuth client from the current TOOL_SERVER_CONNECTIONS
888 config if it hasn't been registered on this node yet.
889 """
890 if client_id in self.clients: 890 ↛ 891line 890 didn't jump to line 891 because the condition on line 890 was never true
891 return self.clients[client_id]['client']
893 try:
894 connections = await Config.get('tool_server.connections', [])
895 except Exception:
896 connections = []
898 for connection in connections or []:
899 if connection.get('type', 'openapi') != 'mcp': 899 ↛ 901line 899 didn't jump to line 901 because the condition on line 899 was always true
900 continue
901 if connection.get('auth_type', 'none') not in ('oauth_2.1', 'oauth_2.1_static'):
902 continue
904 server_id = (connection.get('info') or {}).get('id')
905 if not server_id:
906 continue
908 expected_client_id = f'mcp:{server_id}'
909 if client_id != expected_client_id:
910 continue
912 oauth_client_info = (connection.get('info') or {}).get('oauth_client_info', '')
913 if not oauth_client_info:
914 continue
916 try:
917 oauth_client_info = resolve_oauth_client_info(connection)
918 oauth_client_info = await recover_static_oauth_client_metadata(connection, oauth_client_info)
919 oauth_client_info = apply_connection_oauth_options(connection, oauth_client_info)
920 return self.add_client(expected_client_id, OAuthClientInformationFull(**oauth_client_info))['client']
921 except InvalidToken:
922 log.error(
923 'Failed to lazily add OAuth client %s from config: InvalidToken. '
924 'Stored OAuth client data is invalid; reconnect this tool server.',
925 expected_client_id,
926 )
927 continue
928 except Exception as e:
929 log.error(
930 'Failed to lazily add OAuth client %s from config: %s',
931 expected_client_id,
932 f'{type(e).__name__}: {e}' if str(e) else type(e).__name__,
933 )
934 continue
936 return None
938 def remove_client(self, client_id):
939 if client_id in self.clients:
940 del self.clients[client_id]
941 log.info('Removed OAuth client %s', client_id)
943 if hasattr(self.oauth, '_clients'):
944 if client_id in self.oauth._clients:
945 self.oauth._clients.pop(client_id, None)
947 if hasattr(self.oauth, '_registry'):
948 if client_id in self.oauth._registry:
949 self.oauth._registry.pop(client_id, None)
951 return True
953 async def _preflight_authorization_url(self, client, client_info: OAuthClientInformationFull) -> bool:
954 # TODO: Replace this logic with a more robust OAuth client registration validation
955 # Only perform preflight checks for Starlette OAuth clients
956 if not hasattr(client, 'create_authorization_url'):
957 return True
959 redirect_uri = None
960 if client_info.redirect_uris:
961 redirect_uri = str(client_info.redirect_uris[0])
963 try:
964 kwargs = build_oauth_request_params(client_info)
965 auth_data = await client.create_authorization_url(redirect_uri=redirect_uri, **kwargs)
966 authorization_url = auth_data.get('url')
968 if not authorization_url:
969 return True
970 except Exception as e:
971 log.debug('Skipping OAuth preflight for client %s: %s', client_info.client_id, e)
972 return True
974 try:
975 async with aiohttp.ClientSession(trust_env=True) as session:
976 async with session.get(
977 authorization_url,
978 allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS,
979 ssl=AIOHTTP_CLIENT_SESSION_SSL,
980 ) as resp:
981 if resp.status < 400:
982 return True
983 response_text = await resp.text()
985 error = None
986 error_description = ''
988 content_type = resp.headers.get('content-type', '')
989 if 'application/json' in content_type:
990 try:
991 payload = JSONCodec.loads(response_text)
992 error = payload.get('error')
993 error_description = payload.get('error_description', '')
994 except Exception:
995 pass
996 else:
997 error_description = response_text
999 error_message = f'{error or ""} {error_description or ""}'.lower()
1001 if any(
1002 keyword in error_message
1003 for keyword in (
1004 'invalid_client',
1005 'invalid client',
1006 'client id',
1007 'redirect_uri',
1008 'redirect uri',
1009 )
1010 ):
1011 log.warning(
1012 f'OAuth client preflight detected invalid registration for {client_info.client_id}: {error} {error_description}'
1013 )
1015 return False
1016 except Exception as e:
1017 log.debug('Skipping OAuth preflight network check for client %s: %s', client_info.client_id, e)
1019 return True
1021 async def get_client(self, client_id):
1022 if client_id not in self.clients: 1022 ↛ 1025line 1022 didn't jump to line 1025 because the condition on line 1022 was always true
1023 await self.ensure_client_from_config(client_id)
1025 client = self.clients.get(client_id)
1026 return client['client'] if client else None
1028 async def get_client_info(self, client_id):
1029 if client_id not in self.clients: 1029 ↛ 1032line 1029 didn't jump to line 1032 because the condition on line 1029 was always true
1030 await self.ensure_client_from_config(client_id)
1032 client = self.clients.get(client_id)
1033 return client['client_info'] if client else None
1035 async def get_server_metadata_url(self, client_id):
1036 client = await self.get_client(client_id)
1037 if not client:
1038 return None
1040 return client._server_metadata_url if hasattr(client, '_server_metadata_url') else None
1042 async def get_oauth_token(self, user_id: str, client_id: str, force_refresh: bool = False):
1043 """
1044 Get a valid OAuth token for the user, automatically refreshing if needed.
1046 Args:
1047 user_id: The user ID
1048 client_id: The OAuth client ID (provider)
1049 force_refresh: Force token refresh even if current token appears valid
1051 Returns:
1052 dict: OAuth token data with access_token, or None if no valid token available
1053 """
1054 try:
1055 # Get the OAuth session
1056 session = await OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id)
1057 if not session:
1058 log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}')
1059 return None
1061 if (
1062 force_refresh
1063 or session.expires_at is None
1064 or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at)
1065 ):
1066 log.debug('Token refresh needed for user %s, client_id %s', user_id, session.provider)
1067 refreshed_token = await self._refresh_token(session)
1068 if refreshed_token:
1069 return refreshed_token
1070 else:
1071 log.warning(
1072 f'Token refresh failed for user {user_id}, client_id {session.provider}, deleting session {session.id}'
1073 )
1074 await OAuthSessions.delete_session_by_id(session.id)
1075 return None
1076 return session.token
1078 except Exception as e:
1079 log.error(f'Error getting OAuth token for user {user_id}: {e}')
1080 return None
1082 async def _refresh_token(self, session) -> dict:
1083 """
1084 Refresh an OAuth token if needed, with concurrency protection.
1086 Args:
1087 session: The OAuth session object
1089 Returns:
1090 dict: Refreshed token data, or None if refresh failed
1091 """
1092 try:
1093 # Perform the actual refresh
1094 refreshed_token = await self._perform_token_refresh(session)
1096 if refreshed_token:
1097 # Update the session with new token data
1098 session = await OAuthSessions.update_session_by_id(session.id, refreshed_token)
1099 log.info('Successfully refreshed token for session %s', session.id)
1100 return session.token
1101 else:
1102 log.error(f'Failed to refresh token for session {session.id}')
1103 return None
1105 except Exception as e:
1106 log.error(f'Error refreshing token for session {session.id}: {e}')
1107 return None
1109 async def _perform_token_refresh(self, session) -> dict:
1110 """
1111 Perform the actual OAuth token refresh.
1113 Args:
1114 session: The OAuth session object
1116 Returns:
1117 dict: New token data, or None if refresh failed
1118 """
1119 auth_config = await get_oauth_runtime_config()
1120 client_id = session.provider
1121 token_data = session.token
1123 if not token_data.get('refresh_token'):
1124 log.warning(f'No refresh token available for session {session.id}')
1125 return None
1127 try:
1128 client = await self.get_client(client_id)
1129 if not client:
1130 log.error(f'No OAuth client found for provider {client_id}')
1131 return None
1133 token_endpoint = None
1134 async with aiohttp.ClientSession(trust_env=True) as session_http:
1135 async with session_http.get(await self.get_server_metadata_url(client_id)) as r:
1136 if r.status == 200:
1137 openid_data = await r.json()
1138 token_endpoint = openid_data.get('token_endpoint')
1139 else:
1140 log.error(f'Failed to fetch OpenID configuration for client_id {client_id}')
1141 if not token_endpoint:
1142 log.error(f'No token endpoint found for client_id {client_id}')
1143 return None
1145 # Prepare refresh request
1146 refresh_data = {
1147 'grant_type': 'refresh_token',
1148 'refresh_token': token_data['refresh_token'],
1149 'client_id': client.client_id,
1150 }
1151 client_info = await self.get_client_info(client_id)
1152 if should_send_oauth_resource(client_info):
1153 refresh_data['resource'] = client_info.resource
1155 if hasattr(client, 'client_secret') and client.client_secret:
1156 refresh_data['client_secret'] = client.client_secret
1158 # Add scope if available in client kwargs (some providers require it on refresh)
1159 if (
1160 hasattr(client, 'client_kwargs')
1161 and client.client_kwargs.get('scope')
1162 and auth_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
1163 ):
1164 refresh_data['scope'] = client.client_kwargs['scope']
1166 # Make refresh request
1167 async with aiohttp.ClientSession(trust_env=True) as session_http:
1168 async with session_http.post(
1169 token_endpoint,
1170 data=refresh_data,
1171 headers={'Content-Type': 'application/x-www-form-urlencoded'},
1172 ssl=AIOHTTP_CLIENT_SESSION_SSL,
1173 ) as r:
1174 if r.status == 200:
1175 new_token_data = await r.json()
1177 # Merge with existing token data (preserve refresh_token if not provided)
1178 if 'refresh_token' not in new_token_data:
1179 new_token_data['refresh_token'] = token_data['refresh_token']
1181 _normalize_token_expiry(new_token_data)
1183 log.debug('Token refresh successful for client_id %s', client_id)
1184 return new_token_data
1185 else:
1186 error_text = await r.text()
1187 log.error(f'Token refresh failed for client_id {client_id}: {r.status} - {error_text}')
1188 return None
1190 except Exception as e:
1191 log.error(f'Exception during token refresh for client_id {client_id}: {e}')
1192 return None
1194 async def handle_authorize(self, request, client_id: str, user_id: str) -> RedirectResponse:
1195 client = await self.get_client(client_id)
1196 if client is None:
1197 raise HTTPException(404)
1198 client_info = await self.get_client_info(client_id)
1199 if client_info is None:
1200 # get_client registers client_info too
1201 client_info = await self.get_client_info(client_id)
1202 if client_info is None:
1203 raise HTTPException(404)
1205 redirect_uri = client_info.redirect_uris[0] if client_info.redirect_uris else None
1206 redirect_uri_str = str(redirect_uri) if redirect_uri else None
1207 # Pass explicit scope/resource parameters for providers that require them.
1208 kwargs = build_oauth_request_params(client_info)
1209 try:
1210 auth_data = await client.create_authorization_url(redirect_uri_str, **kwargs)
1211 if not auth_data.get('state'):
1212 raise HTTPException(
1213 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
1214 detail='OAuth authorization state was not generated',
1215 )
1216 auth_data['user_id'] = user_id
1217 await client.save_authorize_data(request, redirect_uri=redirect_uri_str, **auth_data)
1218 return RedirectResponse(auth_data['url'], status_code=302)
1219 except RuntimeError as e:
1220 # authlib raises RuntimeError('Missing "authorize_url" value') when the
1221 # authorization endpoint could not be resolved from server metadata.
1222 # Surface a clear 400 instead of an uncaught 500 for clients that were
1223 # registered before discovery was validated. (#26647)
1224 log.error(f'OAuth authorize failed for client {client_id}: {e}')
1225 raise HTTPException(
1226 status_code=400,
1227 detail=(
1228 'OAuth authorization endpoint could not be resolved for this '
1229 'client. Re-register the MCP server; its OAuth discovery '
1230 'documents may be missing or unreachable.'
1231 ),
1232 )
1234 async def handle_callback(self, request, client_id: str, response):
1235 client = await self.get_client(client_id)
1236 if client is None: 1236 ↛ 1239line 1236 didn't jump to line 1239 because the condition on line 1236 was always true
1237 raise HTTPException(404)
1239 error_message = None
1240 state = request.query_params.get('state')
1241 user_id = None
1242 try:
1243 client_info = await self.get_client_info(client_id)
1244 state_data = await client.framework.get_state_data(request.session, state) if state else None
1245 user_id = state_data.get('user_id') if state_data else None
1246 if not user_id:
1247 raise HTTPException(
1248 status_code=status.HTTP_400_BAD_REQUEST,
1249 detail='OAuth callback state is invalid or expired',
1250 )
1252 if not await get_verified_user_by_id(user_id):
1253 raise HTTPException(
1254 status_code=status.HTTP_401_UNAUTHORIZED,
1255 detail='OAuth callback user is not authorized',
1256 )
1258 request_user = await get_optional_verified_user_from_request(request)
1259 if request_user and request_user.id != user_id:
1260 raise HTTPException(
1261 status_code=status.HTTP_401_UNAUTHORIZED,
1262 detail='OAuth callback user does not match authenticated session',
1263 )
1265 # Note: Do NOT pass client_id/client_secret explicitly here.
1266 # The Authlib client already has these configured during add_client().
1267 # Passing them again causes Authlib to concatenate them (e.g., "ID1,ID1"),
1268 # which results in 401 errors from the token endpoint. (Fix for #19823)
1269 token_kwargs = {}
1270 if should_send_oauth_resource(client_info):
1271 token_kwargs['resource'] = client_info.resource
1272 token = await client.authorize_access_token(request, **token_kwargs)
1274 # Validate that we received a proper token response
1275 # If token exchange failed (e.g., 401), we may get an error response instead
1276 if token and not token.get('access_token'):
1277 error_desc = token.get('error_description', token.get('error', 'Unknown error'))
1278 error_message = f'Token exchange failed: {error_desc}'
1279 log.error('Invalid token response for client_id %s: %s', client_id, error_desc)
1280 token = None
1282 if token:
1283 try:
1284 _normalize_token_expiry(token)
1286 # Clean up any existing sessions for this user/client_id first
1287 sessions = await OAuthSessions.get_sessions_by_user_id(user_id)
1288 for session in sessions:
1289 if session.provider == client_id:
1290 await OAuthSessions.delete_session_by_id(session.id)
1292 session = await OAuthSessions.create_session(
1293 user_id=user_id,
1294 provider=client_id,
1295 token=token,
1296 )
1297 log.info('Stored OAuth session server-side for user %s, client_id %s', user_id, client_id)
1298 except Exception as e:
1299 error_message = 'Failed to store OAuth session server-side'
1300 log.error(f'Failed to store OAuth session server-side: {e}')
1301 else:
1302 if not error_message:
1303 error_message = 'Failed to obtain OAuth token'
1304 log.warning(error_message)
1305 except Exception as e:
1306 error_message = _build_oauth_callback_error_message(e)
1307 log.warning(
1308 'OAuth callback error for user_id=%s client_id=%s: %s',
1309 user_id,
1310 client_id,
1311 error_message,
1312 exc_info=True,
1313 )
1314 finally:
1315 if state and client is not None:
1316 await client.framework.clear_state_data(request.session, state)
1318 webui_url = await Config.get('webui.url')
1319 redirect_url = (str(webui_url or request.base_url)).rstrip('/')
1321 if error_message:
1322 log.debug(error_message)
1323 redirect_url = f'{redirect_url}/?error={urllib.parse.quote_plus(error_message)}'
1324 return RedirectResponse(url=redirect_url, headers=response.headers)
1326 response = RedirectResponse(url=redirect_url, headers=response.headers)
1327 return response
1330class OAuthManager:
1331 def __init__(self, app):
1332 self.oauth = OAuth()
1333 self.app = app
1335 self._clients = {}
1337 for name, provider_config in OAUTH_PROVIDERS.items(): 1337 ↛ 1338line 1337 didn't jump to line 1338 because the loop on line 1337 never started
1338 if 'register' not in provider_config:
1339 log.error(f'OAuth provider {name} missing register function')
1340 continue
1342 client = provider_config['register'](self.oauth)
1343 self._clients[name] = client
1345 def get_client(self, provider_name):
1346 if provider_name not in self._clients:
1347 self._clients[provider_name] = self.oauth.create_client(provider_name)
1348 return self._clients[provider_name]
1350 def get_server_metadata_url(self, provider_name):
1351 if provider_name in self._clients:
1352 client = self._clients[provider_name]
1353 return client._server_metadata_url if hasattr(client, '_server_metadata_url') else None
1354 return None
1356 async def get_oauth_token(self, user_id: str, session_id: str, force_refresh: bool = False):
1357 """
1358 Get a valid OAuth token for the user, automatically refreshing if needed.
1360 Args:
1361 user_id: The user ID
1362 provider: Optional provider name. If None, gets the most recent session.
1363 force_refresh: Force token refresh even if current token appears valid
1365 Returns:
1366 dict: OAuth token data with access_token, or None if no valid token available
1367 """
1368 try:
1369 # Get the OAuth session
1370 session = await OAuthSessions.get_session_by_id_and_user_id(session_id, user_id)
1371 if not session:
1372 log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
1373 return None
1375 # Guard: MCP-provider sessions must be refreshed by
1376 # oauth_client_manager, not the SSO OAuthManager. If one
1377 # reaches here (e.g. via a stale cookie), bail out early
1378 # instead of attempting a refresh that will fail and delete
1379 # the session (#24618).
1380 if (session.provider or '').startswith('mcp:'):
1381 log.debug(
1382 'Skipping MCP session %s (provider=%s) in SSO OAuthManager — handled by oauth_client_manager',
1383 session.id,
1384 session.provider,
1385 )
1386 return None
1388 # SSO integrations may consume the ID token as well as the access token.
1389 expires_at = session.expires_at
1390 id_token = session.token.get('id_token')
1391 if id_token and expires_at is not None:
1392 try:
1393 exp = jwt.decode(id_token, options={'verify_signature': False}).get('exp')
1394 if exp is not None:
1395 expires_at = min(expires_at, int(exp))
1396 except Exception as e:
1397 log.debug('Could not read exp from id_token: %s', e)
1399 if (
1400 force_refresh
1401 or expires_at is None
1402 or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(expires_at)
1403 ):
1404 log.debug('Token refresh needed for user %s, provider %s', user_id, session.provider)
1405 refreshed_token = await self._refresh_token(session)
1406 if refreshed_token:
1407 return refreshed_token
1408 else:
1409 log.warning(
1410 f'Token refresh failed for user {user_id}, provider {session.provider}, deleting session {session.id}'
1411 )
1412 await OAuthSessions.delete_session_by_id(session.id)
1414 return None
1415 return session.token
1417 except Exception as e:
1418 log.error(f'Error getting OAuth token for user {user_id}: {e}')
1419 return None
1421 async def _refresh_token(self, session) -> dict:
1422 """
1423 Refresh an OAuth token if needed, with concurrency protection.
1425 Args:
1426 session: The OAuth session object
1428 Returns:
1429 dict: Refreshed token data, or None if refresh failed
1430 """
1431 try:
1432 # Perform the actual refresh
1433 refreshed_token = await self._perform_token_refresh(session)
1435 if refreshed_token:
1436 # Update the session with new token data
1437 session = await OAuthSessions.update_session_by_id(session.id, refreshed_token)
1438 log.info('Successfully refreshed token for session %s', session.id)
1439 return session.token
1440 else:
1441 log.error(f'Failed to refresh token for session {session.id}')
1442 return None
1444 except Exception as e:
1445 log.error(f'Error refreshing token for session {session.id}: {e}')
1446 return None
1448 async def _perform_token_refresh(self, session) -> dict:
1449 """
1450 Perform the actual OAuth token refresh.
1452 Args:
1453 session: The OAuth session object
1455 Returns:
1456 dict: New token data, or None if refresh failed
1457 """
1458 provider = session.provider
1459 token_data = session.token
1460 auth_config = await get_oauth_runtime_config()
1462 if not token_data.get('refresh_token'):
1463 log.warning(f'No refresh token available for session {session.id}')
1464 return None
1466 try:
1467 client = self.get_client(provider)
1468 if not client:
1469 log.error(f'No OAuth client found for provider {provider}')
1470 return None
1472 server_metadata_url = self.get_server_metadata_url(provider)
1473 token_endpoint = None
1474 async with aiohttp.ClientSession(trust_env=True) as session_http:
1475 async with session_http.get(server_metadata_url) as r:
1476 if r.status == 200:
1477 openid_data = await r.json()
1478 token_endpoint = openid_data.get('token_endpoint')
1479 else:
1480 log.error(f'Failed to fetch OpenID configuration for provider {provider}')
1481 if not token_endpoint:
1482 log.error(f'No token endpoint found for provider {provider}')
1483 return None
1485 # Prepare refresh request
1486 refresh_data = {
1487 'grant_type': 'refresh_token',
1488 'refresh_token': token_data['refresh_token'],
1489 'client_id': client.client_id,
1490 }
1491 # Add client_secret if available (some providers require it)
1492 if hasattr(client, 'client_secret') and client.client_secret:
1493 refresh_data['client_secret'] = client.client_secret
1495 # Add scope if available in client kwargs (some providers require it on refresh)
1496 if (
1497 hasattr(client, 'client_kwargs')
1498 and client.client_kwargs.get('scope')
1499 and auth_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
1500 ):
1501 refresh_data['scope'] = client.client_kwargs['scope']
1503 # Make refresh request
1504 async with aiohttp.ClientSession(trust_env=True) as session_http:
1505 async with session_http.post(
1506 token_endpoint,
1507 data=refresh_data,
1508 headers={'Content-Type': 'application/x-www-form-urlencoded'},
1509 ssl=AIOHTTP_CLIENT_SESSION_SSL,
1510 ) as r:
1511 if r.status == 200:
1512 new_token_data = await r.json()
1514 # Merge with existing token data (preserve refresh_token if not provided)
1515 if 'refresh_token' not in new_token_data:
1516 new_token_data['refresh_token'] = token_data['refresh_token']
1518 _normalize_token_expiry(new_token_data)
1520 log.debug('Token refresh successful for provider %s', provider)
1521 return new_token_data
1522 else:
1523 error_text = await r.text()
1524 log.error(f'Token refresh failed for provider {provider}: {r.status} - {error_text}')
1525 return None
1527 except Exception as e:
1528 log.error(f'Exception during token refresh for provider {provider}: {e}')
1529 return None
1531 async def get_user_role(self, user, user_data, *, access_token: str | None = None):
1532 auth_config = await get_oauth_runtime_config()
1533 user_count = await Users.get_num_users()
1534 if user and user_count == 1:
1535 # If the user is the only user, assign the role "admin" - actually repairs role for single user on login
1536 log.debug('Assigning the only user the admin role')
1537 return 'admin'
1538 if not user and user_count == 0:
1539 # First-user bootstrap: skip role management gating so the
1540 # instance can be initialized. We intentionally return the
1541 # default role here (not 'admin') — admin promotion happens
1542 # race-safely *after* insert via get_num_users() == 1.
1543 log.debug('First user bootstrap: using default role (admin promotion deferred to post-insert)')
1544 return auth_config.DEFAULT_USER_ROLE
1546 if auth_config.ENABLE_OAUTH_ROLE_MANAGEMENT:
1547 log.debug('Running OAUTH Role management')
1548 oauth_claim = auth_config.OAUTH_ROLES_CLAIM
1549 oauth_allowed_roles = auth_config.OAUTH_ALLOWED_ROLES
1550 oauth_admin_roles = auth_config.OAUTH_ADMIN_ROLES
1551 oauth_roles = []
1552 # Keep existing users at their current role unless the provider sent roles.
1553 role = user.role if user else auth_config.DEFAULT_USER_ROLE
1555 if oauth_claim:
1556 claim_data = _get_roles_claim(user_data, oauth_claim)
1557 if claim_data is None and access_token is not None:
1558 # The exchange endpoint has already validated this token with the provider's userinfo endpoint.
1559 try:
1560 token_claims = jwt.decode(access_token, options={'verify_signature': False})
1561 claim_data = _get_roles_claim(token_claims, oauth_claim)
1562 except jwt.PyJWTError as e:
1563 log.debug('Token exchange: cannot decode token claims: %s', e)
1565 if isinstance(claim_data, list):
1566 oauth_roles = claim_data
1567 elif isinstance(claim_data, str):
1568 # Split by the configured separator if present
1569 if OAUTH_ROLES_SEPARATOR and OAUTH_ROLES_SEPARATOR in claim_data:
1570 oauth_roles = claim_data.split(OAUTH_ROLES_SEPARATOR)
1571 else:
1572 oauth_roles = [claim_data]
1573 elif isinstance(claim_data, int):
1574 oauth_roles = [str(claim_data)]
1576 if access_token is not None and not oauth_roles and oauth_allowed_roles and '*' not in oauth_allowed_roles:
1577 log.warning('Token exchange denied: no readable roles claim in userinfo or the token')
1578 raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
1580 log.debug('Oauth Roles claim: %s', oauth_claim)
1581 log.debug('User roles from oauth: %s', oauth_roles)
1582 log.debug('Accepted user roles: %s', oauth_allowed_roles)
1583 log.debug('Accepted admin roles: %s', oauth_admin_roles)
1585 # If roles are present in the token, they must match; otherwise deny access
1586 if oauth_roles:
1587 matched = False
1588 for allowed_role in oauth_allowed_roles:
1589 if allowed_role == '*' or allowed_role in oauth_roles:
1590 log.debug('Assigned user the user role')
1591 role = 'user'
1592 matched = True
1593 break
1594 for admin_role in oauth_admin_roles:
1595 if admin_role in oauth_roles:
1596 log.debug('Assigned user the admin role')
1597 role = 'admin'
1598 matched = True
1599 break
1600 if not matched:
1601 log.warning(
1602 f'OAuth role management enabled but user roles do not match any allowed/admin roles. '
1603 f'User roles: {oauth_roles}, allowed: {oauth_allowed_roles}, admin: {oauth_admin_roles}'
1604 )
1605 raise HTTPException(
1606 status.HTTP_403_FORBIDDEN,
1607 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
1608 )
1609 else:
1610 if not user:
1611 # If role management is disabled, use the default role for new users
1612 role = auth_config.DEFAULT_USER_ROLE
1613 else:
1614 # If role management is disabled, use the existing role for existing users
1615 role = user.role
1617 return role
1619 async def update_user_role_from_oauth(
1620 self,
1621 request,
1622 user,
1623 user_data,
1624 provider,
1625 *,
1626 access_token: str | None = None,
1627 db=None,
1628 ):
1629 determined_role = await self.get_user_role(user, user_data, access_token=access_token)
1630 if user.role == determined_role:
1631 return user
1633 updated_user = await Users.update_user_role_by_id(user.id, determined_role, db=db)
1634 user = updated_user or user
1635 user.role = determined_role
1636 await publish_event(
1637 request,
1638 EVENTS.USER_ROLE_UPDATED,
1639 actor=user,
1640 subject_id=user.id,
1641 source='oauth',
1642 data={'role': determined_role, 'provider': provider},
1643 )
1645 return user
1647 async def update_user_groups(self, request, user, user_data, default_permissions, db=None):
1648 auth_config = await get_oauth_runtime_config()
1649 log.debug('Running OAUTH Group management')
1650 oauth_claim = auth_config.OAUTH_GROUPS_CLAIM
1652 blocked_groups = _parse_blocked_groups(auth_config.OAUTH_BLOCKED_GROUPS)
1654 user_oauth_groups = []
1655 # Nested claim search for groups claim
1656 if oauth_claim:
1657 claim_data = user_data
1658 nested_claims = oauth_claim.split('.')
1659 for nested_claim in nested_claims:
1660 claim_data = claim_data.get(nested_claim, {})
1662 if isinstance(claim_data, list):
1663 user_oauth_groups = claim_data
1664 elif isinstance(claim_data, str): 1664 ↛ anywhereline 1664 didn't jump anywhere: it always raised an exception.
1665 # Split by the configured separator if present
1666 if OAUTH_GROUPS_SEPARATOR in claim_data:
1667 user_oauth_groups = claim_data.split(OAUTH_GROUPS_SEPARATOR)
1668 else:
1669 user_oauth_groups = [claim_data]
1670 else:
1671 user_oauth_groups = []
1673 user_current_groups: list[GroupModel] = await Groups.get_groups_by_member_id(user.id, db=db)
1674 all_available_groups: list[GroupModel] = await Groups.get_all_groups(db=db)
1676 # Create groups if they don't exist and creation is enabled
1677 if auth_config.ENABLE_OAUTH_GROUP_CREATION:
1678 log.debug('Checking for missing groups to create...')
1679 all_group_names = {g.name for g in all_available_groups}
1680 groups_created = False
1681 # Determine creator ID: Prefer admin, fallback to current user if no admin exists
1682 admin_user = await Users.get_super_admin_user()
1683 creator_id = admin_user.id if admin_user else user.id
1684 log.debug('Using creator ID %s for potential group creation.', creator_id)
1686 for group_name in user_oauth_groups:
1687 if group_name not in all_group_names:
1688 log.info("Group '%s' not found via OAuth claim. Creating group...", group_name)
1689 try:
1690 new_group_form = GroupForm(
1691 name=group_name,
1692 description=f"Group '{group_name}' created automatically via OAuth.",
1693 permissions=default_permissions, # Use default permissions from function args
1694 data={'config': {'share': auth_config.OAUTH_GROUP_DEFAULT_SHARE}},
1695 )
1696 # Use determined creator ID (admin or fallback to current user)
1697 created_group = await Groups.insert_new_group(creator_id, new_group_form, db=db)
1698 if created_group:
1699 log.info(
1700 "Successfully created group '%s' with ID %s using creator ID %s",
1701 group_name,
1702 created_group.id,
1703 creator_id,
1704 )
1705 groups_created = True
1706 # Add to local set to prevent duplicate creation attempts in this run
1707 all_group_names.add(group_name)
1708 await publish_event(
1709 request,
1710 EVENTS.GROUP_CREATED,
1711 subject_id=created_group.id,
1712 source='oauth',
1713 data={'name': created_group.name},
1714 )
1715 else:
1716 log.error(f"Failed to create group '{group_name}' via OAuth.")
1717 except Exception as e:
1718 log.error(f"Error creating group '{group_name}' via OAuth: {e}")
1720 # Refresh the list of all available groups if any were created
1721 if groups_created:
1722 all_available_groups = await Groups.get_all_groups(db=db)
1723 log.debug('Refreshed list of all available groups after creation.')
1725 log.debug('Oauth Groups claim: %s', oauth_claim)
1726 log.debug('User oauth groups: %s', user_oauth_groups)
1727 log.debug("User's current groups: %s", [g.name for g in user_current_groups])
1728 log.debug('All groups available in OpenWebUI: %s', [g.name for g in all_available_groups])
1730 # Remove groups that user is no longer a part of
1731 for group_model in user_current_groups:
1732 if (
1733 user_oauth_groups
1734 and group_model.name not in user_oauth_groups
1735 and not is_in_blocked_groups(group_model.name, blocked_groups)
1736 ):
1737 # Remove group from user
1738 log.debug('Removing user from group %s as it is no longer in their oauth groups', group_model.name)
1739 if await Groups.remove_users_from_group(group_model.id, [user.id], db=db):
1740 await publish_event(
1741 request,
1742 EVENTS.GROUP_MEMBER_REMOVED,
1743 actor=user,
1744 subject_id=group_model.id,
1745 source='oauth',
1746 data={'user_ids': [user.id]},
1747 )
1749 # In case a group is created, but perms are never assigned to the group by hitting "save"
1750 group_permissions = group_model.permissions
1751 if not group_permissions:
1752 group_permissions = default_permissions
1754 await Groups.update_group_by_id(
1755 id=group_model.id,
1756 form_data=GroupUpdateForm(
1757 name=group_model.name,
1758 description=group_model.description,
1759 permissions=group_permissions,
1760 ),
1761 overwrite=False,
1762 db=db,
1763 )
1765 # Add user to new groups
1766 for group_model in all_available_groups:
1767 if (
1768 user_oauth_groups
1769 and group_model.name in user_oauth_groups
1770 and not any(gm.name == group_model.name for gm in user_current_groups)
1771 and not is_in_blocked_groups(group_model.name, blocked_groups)
1772 ):
1773 # Add user to group
1774 log.debug('Adding user to group %s as it was found in their oauth groups', group_model.name)
1776 if await Groups.add_users_to_group(group_model.id, [user.id], db=db):
1777 await publish_event(
1778 request,
1779 EVENTS.GROUP_MEMBER_ADDED,
1780 actor=user,
1781 subject_id=group_model.id,
1782 source='oauth',
1783 data={'user_ids': [user.id]},
1784 )
1786 # In case a group is created, but perms are never assigned to the group by hitting "save"
1787 group_permissions = group_model.permissions
1788 if not group_permissions:
1789 group_permissions = default_permissions
1791 await Groups.update_group_by_id(
1792 id=group_model.id,
1793 form_data=GroupUpdateForm(
1794 name=group_model.name,
1795 description=group_model.description,
1796 permissions=group_permissions,
1797 ),
1798 overwrite=False,
1799 db=db,
1800 )
1802 async def _process_picture_url(self, picture_url: str, access_token: str = None) -> str:
1803 """Process a picture URL and return a base64 encoded data URL.
1805 Args:
1806 picture_url: The URL of the picture to process
1807 access_token: Optional OAuth access token for authenticated requests
1809 Returns:
1810 A data URL containing the base64 encoded picture, or "/user.png" if processing fails
1811 """
1812 if not picture_url:
1813 return '/user.png'
1815 try:
1816 await asyncio.to_thread(validate_url, picture_url)
1818 get_kwargs = {}
1819 if access_token:
1820 get_kwargs['headers'] = {
1821 'Authorization': f'Bearer {access_token}',
1822 }
1823 # get_ssrf_safe_session pins the connect-time IP (defeats DNS rebinding); allow_redirects=False keeps validate_url's vet authoritative.
1824 async with get_ssrf_safe_session() as session:
1825 async with session.get(
1826 picture_url,
1827 **get_kwargs,
1828 ssl=AIOHTTP_CLIENT_SESSION_SSL,
1829 allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS,
1830 ) as resp:
1831 if resp.ok:
1832 upstream_mime = (resp.headers.get('Content-Type', '') or '').split(';', 1)[0].strip().lower()
1833 picture = await resp.read()
1834 base64_encoded_picture = base64.b64encode(picture).decode('utf-8')
1835 try:
1836 return validate_image_url(f'data:{upstream_mime};base64,{base64_encoded_picture}')
1837 except ValueError:
1838 log.warning(
1839 f'Rejected OAuth profile picture from {picture_url}: '
1840 f'MIME {upstream_mime!r} is not allowed'
1841 )
1842 return '/user.png'
1843 else:
1844 log.warning(f'Failed to fetch profile picture from {picture_url}')
1845 return '/user.png'
1846 except Exception as e:
1847 log.error(f"Error processing profile picture '{picture_url}': {e}")
1848 return '/user.png'
1850 async def handle_login(self, request, provider):
1851 auth_config = await get_oauth_runtime_config()
1852 if not auth_config.ENABLE_OAUTH:
1853 raise HTTPException(404)
1854 if provider not in OAUTH_PROVIDERS: 1854 ↛ 1857line 1854 didn't jump to line 1857 because the condition on line 1854 was always true
1855 raise HTTPException(404)
1856 # If the provider has a custom redirect URL, use that, otherwise automatically generate one
1857 client = self.get_client(provider)
1858 if client is None:
1859 raise HTTPException(404)
1860 redirect_uri = (client.server_metadata or {}).get('redirect_uri') or request.url_for(
1861 'oauth_login_callback', provider=provider
1862 )
1864 kwargs = {}
1865 if auth_config.OAUTH_AUDIENCE:
1866 kwargs['audience'] = auth_config.OAUTH_AUDIENCE
1867 if OAUTH_AUTHORIZE_PARAMS:
1868 kwargs.update(OAUTH_AUTHORIZE_PARAMS)
1870 return await client.authorize_redirect(request, redirect_uri, **kwargs)
1872 async def handle_callback(self, request, provider, response, db=None):
1873 auth_config = await get_oauth_runtime_config()
1874 if not auth_config.ENABLE_OAUTH:
1875 raise HTTPException(404)
1876 if provider not in OAUTH_PROVIDERS: 1876 ↛ 1879line 1876 didn't jump to line 1879 because the condition on line 1876 was always true
1877 raise HTTPException(404)
1879 error_message = None
1880 try:
1881 client = self.get_client(provider)
1883 auth_params = {}
1885 if client:
1886 if hasattr(client, 'client_id') and OAUTH_ACCESS_TOKEN_REQUEST_INCLUDE_CLIENT_ID:
1887 auth_params['client_id'] = client.client_id
1889 try:
1890 token = await client.authorize_access_token(request, **auth_params)
1891 except BadSignatureError:
1892 # The IdP likely rotated its signing keys and the cached JWKS
1893 # is stale. Evict the cached key set so the next attempt
1894 # fetches fresh keys from the jwks_uri.
1895 log.warning(
1896 'OIDC bad_signature for provider %s — evicting cached JWKS and retrying',
1897 provider,
1898 )
1899 if hasattr(client, 'server_metadata') and isinstance(client.server_metadata, dict):
1900 client.server_metadata.pop('jwks', None)
1901 try:
1902 token = await client.authorize_access_token(request, **auth_params)
1903 except Exception as retry_exc:
1904 detailed_error = _build_oauth_callback_error_message(retry_exc)
1905 log.warning(
1906 'OAuth callback error during authorize_access_token retry for provider %s: %s',
1907 provider,
1908 detailed_error,
1909 exc_info=True,
1910 )
1911 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1912 except Exception as e:
1913 detailed_error = _build_oauth_callback_error_message(e)
1914 log.warning(
1915 'OAuth callback error during authorize_access_token for provider %s: %s',
1916 provider,
1917 detailed_error,
1918 exc_info=True,
1919 )
1920 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1922 # Try to get userinfo from the token first, some providers include it there
1923 user_data: UserInfo = token.get('userinfo')
1924 # Preserve extra claims from the ID token (e.g. roles, groups for
1925 # Microsoft Entra ID) before the userinfo endpoint possibly overwrites them.
1926 id_token_claims = dict(user_data) if user_data else {}
1927 if (
1928 (not user_data)
1929 or (auth_config.OAUTH_EMAIL_CLAIM not in user_data)
1930 or (auth_config.OAUTH_USERNAME_CLAIM not in user_data)
1931 ):
1932 user_data: UserInfo = await client.userinfo(token=token)
1933 # Merge back ID token claims that the userinfo endpoint doesn't
1934 # return. Only backfill missing keys so userinfo always wins.
1935 if user_data and id_token_claims:
1936 for key, value in id_token_claims.items():
1937 if key not in user_data:
1938 user_data[key] = value
1939 if provider == 'feishu' and isinstance(user_data, dict) and 'data' in user_data:
1940 user_data = user_data['data']
1941 if not user_data:
1942 log.warning('OAuth callback failed for provider %s, user data is missing', provider)
1943 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1945 # Extract the "sub" claim, using custom claim if configured
1946 if auth_config.OAUTH_SUB_CLAIM:
1947 sub = user_data.get(auth_config.OAUTH_SUB_CLAIM)
1948 else:
1949 # Fallback to the default sub claim if not configured
1950 sub = user_data.get(OAUTH_PROVIDERS[provider].get('sub_claim', 'sub'))
1951 if not sub:
1952 log.warning(f'OAuth callback failed, sub is missing: {user_data}')
1953 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1954 sub = str(sub)
1956 oauth_data = {}
1957 oauth_data[provider] = {
1958 'sub': sub,
1959 }
1961 # Email extraction
1962 email_claim = auth_config.OAUTH_EMAIL_CLAIM
1963 email = user_data.get(email_claim, '')
1964 # We currently mandate that email addresses are provided
1965 if not email:
1966 # If the provider is GitHub,and public email is not provided, we can use the access token to fetch the user's email
1967 if provider == 'github':
1968 try:
1969 access_token = token.get('access_token')
1970 headers = {'Authorization': f'Bearer {access_token}'}
1971 async with aiohttp.ClientSession(trust_env=True) as session:
1972 async with session.get(
1973 'https://api.github.com/user/emails',
1974 headers=headers,
1975 ssl=AIOHTTP_CLIENT_SESSION_SSL,
1976 ) as resp:
1977 if resp.ok:
1978 emails = await resp.json()
1979 # use the primary email as the user's email
1980 primary_email = next(
1981 (e['email'] for e in emails if e.get('primary')),
1982 None,
1983 )
1984 if primary_email:
1985 email = primary_email
1986 else:
1987 log.warning('No primary email found in GitHub response')
1988 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1989 else:
1990 log.warning('Failed to fetch GitHub email')
1991 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1992 except Exception as e:
1993 log.warning(f'Error fetching GitHub email: {e}')
1994 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
1995 elif ENABLE_OAUTH_EMAIL_FALLBACK:
1996 email = f'{provider}@{sub}.local'
1997 else:
1998 log.warning(f'OAuth callback failed, email is missing: {user_data}')
1999 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
2001 email = email.lower()
2002 # If allowed domains are configured, check if the email domain is in the list
2003 if (
2004 '*' not in auth_config.OAUTH_ALLOWED_DOMAINS
2005 and email.split('@')[-1] not in auth_config.OAUTH_ALLOWED_DOMAINS
2006 ):
2007 log.warning(f'OAuth callback failed, e-mail domain is not in the list of allowed domains: {user_data}')
2008 raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
2010 # Check if the user exists
2011 user = await Users.get_user_by_oauth_sub(provider, sub, db=db)
2012 if not user:
2013 # If the user does not exist, check if merging is enabled
2014 if auth_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL:
2015 # Check if the user exists by email
2016 user = await Users.get_user_by_email(email, db=db)
2017 if user:
2018 # Update the user with the new oauth sub
2019 user = await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) or user
2021 if user:
2022 provider_oauth = (user.oauth or {}).get(provider) if isinstance(user.oauth, dict) else None
2023 # Lazy repair for legacy rows that stored numeric provider ids as JSON numbers.
2024 if isinstance(provider_oauth, dict) and provider_oauth.get('sub') != sub:
2025 user = await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) or user
2027 if user:
2028 user = await self.update_user_role_from_oauth(
2029 request=request,
2030 user=user,
2031 user_data=user_data,
2032 provider=provider,
2033 db=db,
2034 )
2036 updated_fields = []
2038 if auth_config.OAUTH_UPDATE_NAME_ON_LOGIN:
2039 username_claim = auth_config.OAUTH_USERNAME_CLAIM
2040 if username_claim:
2041 new_name = user_data.get(username_claim)
2042 if new_name and new_name != user.name:
2043 updated_user = await Users.update_user_by_id(user.id, {'name': new_name}, db=db)
2044 if updated_user:
2045 user = updated_user
2046 updated_fields.append('name')
2047 log.debug('Updated name for user %s', user.email)
2049 if auth_config.OAUTH_UPDATE_EMAIL_ON_LOGIN:
2050 email_claim = auth_config.OAUTH_EMAIL_CLAIM
2051 if email_claim:
2052 new_email = user_data.get(email_claim)
2053 if new_email and new_email.lower() != user.email.lower():
2054 existing_user = await Users.get_user_by_email(new_email, db=db)
2055 if existing_user:
2056 log.error(
2057 f'Cannot update email to {new_email} for user {user.id} because it is already taken.'
2058 )
2059 elif await Auths.update_email_by_id(user.id, new_email.lower(), db=db):
2060 user = await Users.get_user_by_id(user.id, db=db) or user
2061 updated_fields.append('email')
2062 log.debug('Updated email for user %s', user.id)
2064 # Update profile picture if enabled and different from current
2065 if auth_config.OAUTH_UPDATE_PICTURE_ON_LOGIN:
2066 picture_claim = auth_config.OAUTH_PICTURE_CLAIM
2067 if picture_claim:
2068 new_picture_url = user_data.get(
2069 picture_claim,
2070 OAUTH_PROVIDERS[provider].get('picture_url', ''),
2071 )
2072 processed_picture_url = await self._process_picture_url(
2073 new_picture_url, token.get('access_token')
2074 )
2075 if processed_picture_url != user.profile_image_url:
2076 updated_user = await Users.update_user_profile_image_url_by_id(
2077 user.id, processed_picture_url, db=db
2078 )
2079 if updated_user:
2080 user = updated_user
2081 updated_fields.append('profile_image_url')
2082 log.debug('Updated profile picture for user %s', user.email)
2084 if updated_fields:
2085 await publish_event(
2086 request,
2087 EVENTS.USER_UPDATED,
2088 actor=user,
2089 subject_id=user.id,
2090 source='oauth',
2091 data={'updated_fields': updated_fields, 'provider': provider},
2092 )
2093 else:
2094 # If the user does not exist, check if signups are enabled
2095 if auth_config.ENABLE_OAUTH_SIGNUP:
2096 # Check if an existing user with the same email already exists
2097 existing_user = await Users.get_user_by_email(email, db=db)
2098 if existing_user:
2099 raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
2101 picture_claim = auth_config.OAUTH_PICTURE_CLAIM
2102 if picture_claim:
2103 picture_url = user_data.get(
2104 picture_claim,
2105 OAUTH_PROVIDERS[provider].get('picture_url', ''),
2106 )
2107 picture_url = await self._process_picture_url(picture_url, token.get('access_token'))
2108 else:
2109 picture_url = '/user.png'
2110 username_claim = auth_config.OAUTH_USERNAME_CLAIM
2112 name = user_data.get(username_claim)
2113 if not name:
2114 log.warning('Username claim is missing, using email as name')
2115 name = email
2117 user = await Auths.insert_new_auth(
2118 email=email,
2119 password=await get_password_hash(str(uuid.uuid4())), # Random password, not used
2120 name=name,
2121 profile_image_url=picture_url,
2122 role=await self.get_user_role(None, user_data),
2123 oauth=oauth_data,
2124 db=db,
2125 )
2127 if not user:
2128 raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR)
2130 # Atomically check if this is the only user *after* the
2131 # insert to avoid TOCTOU race on first-user registration.
2132 # Matches signup_handler pattern.
2133 if await Users.get_num_users(db=db) == 1:
2134 await Users.update_user_role_by_id(user.id, 'admin', db=db)
2135 user = await Users.get_user_by_id(user.id, db=db)
2137 default_group_id = await Config.get('ui.default_group_id')
2138 await apply_default_group_assignment(default_group_id, user.id, db=db)
2139 await publish_event(
2140 request,
2141 EVENTS.USER_CREATED,
2142 actor=user,
2143 subject_id=user.id,
2144 source='oauth',
2145 data={'role': user.role, 'provider': provider},
2146 )
2148 else:
2149 raise HTTPException(
2150 status.HTTP_403_FORBIDDEN,
2151 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
2152 )
2154 jwt_token = create_token(
2155 data={'id': user.id},
2156 expires_delta=parse_duration(auth_config.JWT_EXPIRES_IN),
2157 )
2158 if auth_config.ENABLE_OAUTH_GROUP_MANAGEMENT:
2159 await self.update_user_groups(
2160 request=request,
2161 user=user,
2162 user_data=user_data,
2163 default_permissions=await Config.get('user.permissions'),
2164 db=db,
2165 )
2167 except Exception as e:
2168 log.error(f'Error during OAuth process: {e}')
2169 error_message = (
2170 e.detail
2171 if isinstance(e, HTTPException) and e.detail
2172 else ERROR_MESSAGES.DEFAULT('Error during OAuth process')
2173 )
2175 webui_url = await Config.get('webui.url')
2176 redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
2177 redirect_url = f'{redirect_base_url}/auth'
2179 if error_message:
2180 redirect_url = f'{redirect_url}?error={urllib.parse.quote_plus(error_message)}'
2181 return RedirectResponse(url=redirect_url, headers=response.headers)
2183 response = RedirectResponse(url=redirect_url, headers=response.headers)
2185 # Compute cookie expiry from JWT lifetime
2186 expires_delta = parse_duration(auth_config.JWT_EXPIRES_IN)
2187 cookie_max_age = int(expires_delta.total_seconds()) if expires_delta else None
2189 # Set the cookie token
2190 # Redirect back to the frontend with the JWT token
2191 response.set_cookie(
2192 key='token',
2193 value=jwt_token,
2194 httponly=False, # Required for frontend access
2195 samesite=WEBUI_AUTH_COOKIE_SAME_SITE,
2196 secure=WEBUI_AUTH_COOKIE_SECURE,
2197 **({'max_age': cookie_max_age} if cookie_max_age is not None else {}),
2198 )
2200 await publish_event(
2201 request,
2202 EVENTS.AUTH_LOGIN,
2203 actor=user,
2204 subject_id=user.id,
2205 subject_type='user',
2206 source='oauth',
2207 data={'auth_method': 'oauth', 'provider': provider},
2208 )
2210 # Legacy cookies for compatibility with older frontend versions
2211 if ENABLE_OAUTH_ID_TOKEN_COOKIE:
2212 response.set_cookie(
2213 key='oauth_id_token',
2214 value=token.get('id_token'),
2215 httponly=True,
2216 samesite=WEBUI_AUTH_COOKIE_SAME_SITE,
2217 secure=WEBUI_AUTH_COOKIE_SECURE,
2218 **({'max_age': cookie_max_age} if cookie_max_age is not None else {}),
2219 )
2221 try:
2222 _normalize_token_expiry(token)
2224 # Enforce max concurrent sessions per user/provider to prevent
2225 # unbounded growth while allowing multi-device usage
2226 sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db)
2227 provider_sessions = sorted(
2228 [session for session in sessions if session.provider == provider],
2229 key=lambda session: session.created_at,
2230 reverse=True,
2231 )
2232 # Keep the newest sessions up to the limit, prune the rest
2233 if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER:
2234 for old_session in provider_sessions[OAUTH_MAX_SESSIONS_PER_USER - 1 :]:
2235 await OAuthSessions.delete_session_by_id(old_session.id, db=db)
2237 session = await OAuthSessions.create_session(
2238 user_id=user.id,
2239 provider=provider,
2240 token=token,
2241 db=db,
2242 )
2244 if session:
2245 response.set_cookie(
2246 key='oauth_session_id',
2247 value=session.id,
2248 httponly=True,
2249 samesite=WEBUI_AUTH_COOKIE_SAME_SITE,
2250 secure=WEBUI_AUTH_COOKIE_SECURE,
2251 **({'max_age': cookie_max_age} if cookie_max_age is not None else {}),
2252 )
2254 log.info('Stored OAuth session server-side for user %s, provider %s', user.id, provider)
2255 else:
2256 log.warning(f'Failed to create OAuth session for user {user.id}, provider {provider}')
2257 except Exception as e:
2258 log.error(f'Failed to store OAuth session server-side: {e}')
2260 return response
2262 async def handle_backchannel_logout(self, request, db=None):
2263 """
2264 Handle an OIDC Back-Channel Logout request.
2265 Validates the logout_token, identifies the user, revokes their
2266 sessions via Redis, and deletes their OAuth sessions.
2267 Returns a JSONResponse per the OIDC Back-Channel Logout 1.0 spec.
2268 """
2269 from fastapi.responses import JSONResponse
2271 # 1. Extract logout_token from form body
2272 try:
2273 form = await request.form()
2274 logout_token = form.get('logout_token')
2275 except Exception:
2276 logout_token = None
2278 if not logout_token:
2279 return JSONResponse(
2280 status_code=400,
2281 content={'error': 'invalid_request', 'error_description': 'Missing logout_token parameter'},
2282 )
2284 # 2. Peek at unverified issuer to match against configured providers
2285 try:
2286 unverified_claims = jwt.decode(logout_token, options={'verify_signature': False})
2287 token_issuer = unverified_claims.get('iss')
2288 except Exception as e:
2289 log.warning(f'Back-channel logout: cannot decode logout_token: {e}')
2290 return JSONResponse(
2291 status_code=400,
2292 content={'error': 'invalid_request', 'error_description': 'Malformed logout_token'},
2293 )
2295 if not token_issuer:
2296 return JSONResponse(
2297 status_code=400,
2298 content={'error': 'invalid_request', 'error_description': 'logout_token missing iss claim'},
2299 )
2301 # 3. Find the configured provider whose issuer matches the token
2302 matched_provider = None
2303 matched_client = None
2304 matched_jwks_uri = None
2306 for provider_name in OAUTH_PROVIDERS:
2307 client = self.get_client(provider_name)
2308 if not client:
2309 continue
2311 try:
2312 oidc_config = await client.load_server_metadata()
2313 except Exception as e:
2314 log.debug('Back-channel logout: error checking provider %s: %s', provider_name, e)
2315 continue
2317 if oidc_config.get('issuer') == token_issuer:
2318 matched_provider = provider_name
2319 matched_client = client
2320 matched_jwks_uri = oidc_config.get('jwks_uri')
2321 break
2323 if not matched_provider or not matched_client or not matched_client.client_id or not matched_jwks_uri:
2324 log.warning(f'Back-channel logout: no configured provider matches issuer {token_issuer}')
2325 return JSONResponse(
2326 status_code=400,
2327 content={
2328 'error': 'invalid_request',
2329 'error_description': 'No configured provider matches token issuer',
2330 },
2331 )
2333 # 4. Validate the logout_token signature and claims
2334 try:
2335 token_kid = jwt.get_unverified_header(logout_token).get('kid')
2336 if not token_kid:
2337 raise jwt.InvalidTokenError('logout_token missing kid header')
2339 try:
2340 jwk_set = jwt.PyJWKSet.from_dict(await matched_client.fetch_jwk_set())
2341 except jwt.PyJWTError as e:
2342 raise jwt.InvalidTokenError(str(e))
2344 signing_key = next(
2345 (key for key in jwk_set.keys if key.key_id == token_kid and key.public_key_use in ['sig', None]),
2346 None,
2347 )
2348 if not signing_key:
2349 raise jwt.InvalidTokenError('no signing key matches the token kid')
2351 claims = jwt.decode(
2352 logout_token,
2353 signing_key.key,
2354 algorithms=['RS256', 'RS384', 'RS512', 'ES256', 'ES384', 'ES512'],
2355 audience=matched_client.client_id,
2356 issuer=token_issuer,
2357 options={
2358 'require': ['iss', 'aud', 'iat', 'events'],
2359 },
2360 )
2361 except jwt.InvalidTokenError as e:
2362 log.warning(f'Back-channel logout: invalid logout_token: {e}')
2363 return JSONResponse(
2364 status_code=400,
2365 content={'error': 'invalid_request', 'error_description': f'Invalid logout_token: {e}'},
2366 )
2367 except Exception as e:
2368 log.error(f'Back-channel logout: error validating logout_token: {e}')
2369 return JSONResponse(
2370 status_code=400,
2371 content={'error': 'invalid_request', 'error_description': 'Failed to validate logout_token'},
2372 )
2374 # 5. Validate events claim per spec
2375 events = claims.get('events', {})
2376 if 'http://schemas.openid.net/event/backchannel-logout' not in events:
2377 log.warning('Back-channel logout: missing required backchannel-logout event claim')
2378 return JSONResponse(
2379 status_code=400,
2380 content={'error': 'invalid_request', 'error_description': 'Missing backchannel-logout event claim'},
2381 )
2383 # 6. Per spec, back-channel logout tokens MUST NOT contain a nonce
2384 if 'nonce' in claims:
2385 log.warning('Back-channel logout: logout_token contains nonce (rejected per spec)')
2386 return JSONResponse(
2387 status_code=400,
2388 content={'error': 'invalid_request', 'error_description': 'logout_token must not contain nonce'},
2389 )
2391 # 7. Extract sub and/or sid — at least one must be present
2392 sub = claims.get('sub')
2393 sid = claims.get('sid')
2395 if not sub and not sid:
2396 log.warning('Back-channel logout: logout_token contains neither sub nor sid')
2397 return JSONResponse(
2398 status_code=400,
2399 content={'error': 'invalid_request', 'error_description': 'logout_token must contain sub or sid'},
2400 )
2402 # 8. Identify users to log out
2403 users_to_logout = []
2404 if sub:
2405 user = await Users.get_user_by_oauth_sub(matched_provider, str(sub), db=db)
2406 if user:
2407 users_to_logout.append(user)
2409 if not users_to_logout and sid:
2410 log.debug('Back-channel logout: no user found by sub, sid-based lookup not yet supported (sid=%s)', sid)
2412 if not users_to_logout:
2413 log.debug(
2414 'Back-channel logout: no matching user for provider=%s, sub=%s, sid=%s', matched_provider, sub, sid
2415 )
2416 return JSONResponse(status_code=200, content={})
2418 # 9. Revoke tokens and delete sessions
2419 redis = request.app.state.redis
2420 if not redis:
2421 log.warning(
2422 'Back-channel logout: Redis not configured, cannot revoke JWT tokens. '
2423 'OAuth sessions will be deleted but existing JWTs will remain valid until expiry.'
2424 )
2426 revoked_count = 0
2427 for user in users_to_logout:
2428 sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db)
2429 for oauth_session in sessions:
2430 await OAuthSessions.delete_session_by_id(oauth_session.id, db=db)
2432 if redis:
2433 await revoke_user_tokens(request, user.id)
2434 revoked_count += 1
2436 log.info(
2437 'Back-channel logout: revoked sessions for user %s (email=%s, provider=%s, sessions_deleted=%s)',
2438 user.id,
2439 user.email,
2440 matched_provider,
2441 len(sessions),
2442 )
2444 log.info(
2445 'Back-channel logout: completed for %s user(s), %s revocation(s) set', len(users_to_logout), revoked_count
2446 )
2447 return JSONResponse(status_code=200, content={})