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

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 

15 

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 

96 

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) 

100 

101 

102class OAuthClientMetadata(MCPOAuthClientMetadata): 

103 token_endpoint_auth_method: Literal['none', 'client_secret_basic', 'client_secret_post'] = 'client_secret_post' 

104 pass 

105 

106 

107OAuthResourceParameterMode = Literal['auto', 'include', 'omit'] 

108 

109 

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' 

114 

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 

119 

120 server_metadata: Optional[OAuthMetadata] = None # Fetched from the OAuth server 

121 

122 

123from open_webui.env import GLOBAL_LOG_LEVEL 

124from open_webui.utils.json_codec import JSONCodec 

125 

126logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) 

127log = logging.getLogger(__name__) 

128 

129OAUTH_RESOURCE_PARAMETER_MODES = {'auto', 'include', 'omit'} 

130 

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} 

185 

186 

187def _default_value(value): 

188 return getattr(value, 'value', value) 

189 

190 

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 

199 

200 

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) 

206 

207 

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 

212 

213 

214def _normalize_token_expiry(token: dict) -> dict: 

215 """Ensure a token dict always has a numeric access-token ``expires_at``. 

216 

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. 

225 

226 Also stamps *issued_at* for auditing. 

227 """ 

228 token['issued_at'] = datetime.now().timestamp() 

229 

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 

245 

246 token['expires_at'] = expires_at 

247 return token 

248 

249 

250FERNET = None 

251 

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() 

257 

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 

263 

264 

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 

274 

275 

276def decrypt_data(data: str): 

277 """Decrypt data from storage""" 

278 decrypted = FERNET.decrypt(data.encode()).decode() 

279 return JSONCodec.loads(decrypted) 

280 

281 

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) 

304 

305 detail = detail.replace('\n', ' ').strip() 

306 if not detail: 

307 detail = e.__class__.__name__ 

308 

309 message = f'OAuth callback failed: {detail}' 

310 return message[:197] + '...' if len(message) > 200 else message 

311 

312 

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. 

317 

318 Args: 

319 group_name: The group name to check 

320 groups: List of patterns to match against 

321 

322 Returns: 

323 True if the group is blocked, False otherwise 

324 """ 

325 if not groups: 

326 return False 

327 

328 for group_pattern in groups: 

329 if not group_pattern: # Skip empty patterns 

330 continue 

331 

332 # Exact match 

333 if group_name == group_pattern: 

334 return True 

335 

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 

345 

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 

350 

351 return False 

352 

353 

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] 

365 

366 

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 

371 

372 

373@dataclass 

374class ProtectedResourceMetadata: 

375 """RFC 9728 Protected Resource Metadata fields relevant to OAuth flows.""" 

376 

377 resource: str | None = None 

378 authorization_servers: list[str] = field(default_factory=list) 

379 scopes_supported: list[str] = field(default_factory=list) 

380 

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 

388 

389 

390async def get_protected_resource_metadata(server_url: str) -> ProtectedResourceMetadata: 

391 """ 

392 Fetch RFC 9728 Protected Resource Metadata from an MCP server. 

393 

394 https://modelcontextprotocol.io/specification/2025-03-26/basic/authorization 

395 

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) 

436 

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() 

445 

446 resource = resource_metadata.get('resource') or None 

447 if resource: 

448 log.debug('Discovered resource indicator: %s', resource) 

449 

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) 

454 

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) 

464 

465 return ProtectedResourceMetadata( 

466 resource=resource, authorization_servers=authorization_servers, scopes_supported=scopes 

467 ) 

468 

469 

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 = [] 

474 

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 ) 

484 

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 ) 

491 

492 return urls 

493 

494 

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) 

499 

500 

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 

513 

514 webui_url = await Config.get('webui.url') 

515 redirect_base_url = (str(webui_url or request.base_url)).rstrip('/') 

516 

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 ) 

526 

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 

530 

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) 

539 

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) 

555 

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 ) 

565 

566 break 

567 except Exception as e: 

568 log.error(f'Error parsing OAuth metadata from {url}: {e}') 

569 continue 

570 

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 ) 

584 

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') 

591 

592 registration_data = oauth_client_metadata.model_dump( 

593 exclude_none=True, 

594 mode='json', 

595 by_alias=True, 

596 ) 

597 

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() 

605 

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 

634 

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 

645 

646 

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 

663 

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' 

667 

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 

683 

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 ) 

691 

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] 

700 

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 ) 

713 

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 

721 

722 

723def resolve_oauth_client_info(connection: dict) -> dict: 

724 """ 

725 Decrypt OAuth client info from a tool server connection config. 

726 

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', '')) 

732 

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

737 

738 return data 

739 

740 

741def normalize_oauth_resource_parameter(value: str | None) -> OAuthResourceParameterMode: 

742 if value in OAUTH_RESOURCE_PARAMETER_MODES: 

743 return value 

744 return 'auto' 

745 

746 

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 ) 

753 

754 

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 

760 

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 

768 

769 

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()) 

774 

775 

776def should_send_oauth_resource(client_info: OAuthClientInformationFull | None) -> bool: 

777 if not client_info or not client_info.resource: 

778 return False 

779 

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 

785 

786 return not scope_has_resource_indicator(client_info.scope) 

787 

788 

789def build_oauth_request_params(client_info: OAuthClientInformationFull | None) -> dict: 

790 if not client_info: 

791 return {} 

792 

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 

799 

800 

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 

804 

805 if oauth_client_info.get('scope') and oauth_client_info.get('resource'): 

806 return oauth_client_info 

807 

808 server_url = connection.get('url') 

809 if not server_url: 

810 return oauth_client_info 

811 

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 

817 

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) 

822 

823 if not recovered.get('resource') and resource_metadata.resource: 

824 recovered['resource'] = resource_metadata.resource 

825 

826 return recovered 

827 

828 

829class OAuthClientManager: 

830 def __init__(self, app): 

831 self.oauth = OAuth() 

832 self.app = app 

833 self.clients = {} 

834 

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 } 

852 

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) 

863 

864 # Default to S256 for OAuth 2.1 (PKCE is mandatory per RFC 9700) 

865 kwargs['code_challenge_method'] = 'S256' 

866 

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

878 

879 self.clients[client_id] = { 

880 'client': self.oauth.register(**kwargs), 

881 'client_info': oauth_client_info, 

882 } 

883 return self.clients[client_id] 

884 

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

892 

893 try: 

894 connections = await Config.get('tool_server.connections', []) 

895 except Exception: 

896 connections = [] 

897 

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 

903 

904 server_id = (connection.get('info') or {}).get('id') 

905 if not server_id: 

906 continue 

907 

908 expected_client_id = f'mcp:{server_id}' 

909 if client_id != expected_client_id: 

910 continue 

911 

912 oauth_client_info = (connection.get('info') or {}).get('oauth_client_info', '') 

913 if not oauth_client_info: 

914 continue 

915 

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 

935 

936 return None 

937 

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) 

942 

943 if hasattr(self.oauth, '_clients'): 

944 if client_id in self.oauth._clients: 

945 self.oauth._clients.pop(client_id, None) 

946 

947 if hasattr(self.oauth, '_registry'): 

948 if client_id in self.oauth._registry: 

949 self.oauth._registry.pop(client_id, None) 

950 

951 return True 

952 

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 

958 

959 redirect_uri = None 

960 if client_info.redirect_uris: 

961 redirect_uri = str(client_info.redirect_uris[0]) 

962 

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') 

967 

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 

973 

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() 

984 

985 error = None 

986 error_description = '' 

987 

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 

998 

999 error_message = f'{error or ""} {error_description or ""}'.lower() 

1000 

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 ) 

1014 

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) 

1018 

1019 return True 

1020 

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) 

1024 

1025 client = self.clients.get(client_id) 

1026 return client['client'] if client else None 

1027 

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) 

1031 

1032 client = self.clients.get(client_id) 

1033 return client['client_info'] if client else None 

1034 

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 

1039 

1040 return client._server_metadata_url if hasattr(client, '_server_metadata_url') else None 

1041 

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. 

1045 

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 

1050 

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 

1060 

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 

1077 

1078 except Exception as e: 

1079 log.error(f'Error getting OAuth token for user {user_id}: {e}') 

1080 return None 

1081 

1082 async def _refresh_token(self, session) -> dict: 

1083 """ 

1084 Refresh an OAuth token if needed, with concurrency protection. 

1085 

1086 Args: 

1087 session: The OAuth session object 

1088 

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) 

1095 

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 

1104 

1105 except Exception as e: 

1106 log.error(f'Error refreshing token for session {session.id}: {e}') 

1107 return None 

1108 

1109 async def _perform_token_refresh(self, session) -> dict: 

1110 """ 

1111 Perform the actual OAuth token refresh. 

1112 

1113 Args: 

1114 session: The OAuth session object 

1115 

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 

1122 

1123 if not token_data.get('refresh_token'): 

1124 log.warning(f'No refresh token available for session {session.id}') 

1125 return None 

1126 

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 

1132 

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 

1144 

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 

1154 

1155 if hasattr(client, 'client_secret') and client.client_secret: 

1156 refresh_data['client_secret'] = client.client_secret 

1157 

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

1165 

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() 

1176 

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

1180 

1181 _normalize_token_expiry(new_token_data) 

1182 

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 

1189 

1190 except Exception as e: 

1191 log.error(f'Exception during token refresh for client_id {client_id}: {e}') 

1192 return None 

1193 

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) 

1204 

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 ) 

1233 

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) 

1238 

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 ) 

1251 

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 ) 

1257 

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 ) 

1264 

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) 

1273 

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 

1281 

1282 if token: 

1283 try: 

1284 _normalize_token_expiry(token) 

1285 

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) 

1291 

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) 

1317 

1318 webui_url = await Config.get('webui.url') 

1319 redirect_url = (str(webui_url or request.base_url)).rstrip('/') 

1320 

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) 

1325 

1326 response = RedirectResponse(url=redirect_url, headers=response.headers) 

1327 return response 

1328 

1329 

1330class OAuthManager: 

1331 def __init__(self, app): 

1332 self.oauth = OAuth() 

1333 self.app = app 

1334 

1335 self._clients = {} 

1336 

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 

1341 

1342 client = provider_config['register'](self.oauth) 

1343 self._clients[name] = client 

1344 

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] 

1349 

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 

1355 

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. 

1359 

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 

1364 

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 

1374 

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 

1387 

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) 

1398 

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) 

1413 

1414 return None 

1415 return session.token 

1416 

1417 except Exception as e: 

1418 log.error(f'Error getting OAuth token for user {user_id}: {e}') 

1419 return None 

1420 

1421 async def _refresh_token(self, session) -> dict: 

1422 """ 

1423 Refresh an OAuth token if needed, with concurrency protection. 

1424 

1425 Args: 

1426 session: The OAuth session object 

1427 

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) 

1434 

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 

1443 

1444 except Exception as e: 

1445 log.error(f'Error refreshing token for session {session.id}: {e}') 

1446 return None 

1447 

1448 async def _perform_token_refresh(self, session) -> dict: 

1449 """ 

1450 Perform the actual OAuth token refresh. 

1451 

1452 Args: 

1453 session: The OAuth session object 

1454 

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() 

1461 

1462 if not token_data.get('refresh_token'): 

1463 log.warning(f'No refresh token available for session {session.id}') 

1464 return None 

1465 

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 

1471 

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 

1484 

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 

1494 

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

1502 

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() 

1513 

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

1517 

1518 _normalize_token_expiry(new_token_data) 

1519 

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 

1526 

1527 except Exception as e: 

1528 log.error(f'Exception during token refresh for provider {provider}: {e}') 

1529 return None 

1530 

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 

1545 

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 

1554 

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) 

1564 

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)] 

1575 

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) 

1579 

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) 

1584 

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 

1616 

1617 return role 

1618 

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 

1632 

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 ) 

1644 

1645 return user 

1646 

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 

1651 

1652 blocked_groups = _parse_blocked_groups(auth_config.OAUTH_BLOCKED_GROUPS) 

1653 

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, {}) 

1661 

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 = [] 

1672 

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) 

1675 

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) 

1685 

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}") 

1719 

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.') 

1724 

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]) 

1729 

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 ) 

1748 

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 

1753 

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 ) 

1764 

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) 

1775 

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 ) 

1785 

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 

1790 

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 ) 

1801 

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. 

1804 

1805 Args: 

1806 picture_url: The URL of the picture to process 

1807 access_token: Optional OAuth access token for authenticated requests 

1808 

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' 

1814 

1815 try: 

1816 await asyncio.to_thread(validate_url, picture_url) 

1817 

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' 

1849 

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 ) 

1863 

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) 

1869 

1870 return await client.authorize_redirect(request, redirect_uri, **kwargs) 

1871 

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) 

1878 

1879 error_message = None 

1880 try: 

1881 client = self.get_client(provider) 

1882 

1883 auth_params = {} 

1884 

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 

1888 

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) 

1921 

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) 

1944 

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) 

1955 

1956 oauth_data = {} 

1957 oauth_data[provider] = { 

1958 'sub': sub, 

1959 } 

1960 

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) 

2000 

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) 

2009 

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 

2020 

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 

2026 

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 ) 

2035 

2036 updated_fields = [] 

2037 

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) 

2048 

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) 

2063 

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) 

2083 

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) 

2100 

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 

2111 

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 

2116 

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 ) 

2126 

2127 if not user: 

2128 raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) 

2129 

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) 

2136 

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 ) 

2147 

2148 else: 

2149 raise HTTPException( 

2150 status.HTTP_403_FORBIDDEN, 

2151 detail=ERROR_MESSAGES.ACCESS_PROHIBITED, 

2152 ) 

2153 

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 ) 

2166 

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 ) 

2174 

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' 

2178 

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) 

2182 

2183 response = RedirectResponse(url=redirect_url, headers=response.headers) 

2184 

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 

2188 

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 ) 

2199 

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 ) 

2209 

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 ) 

2220 

2221 try: 

2222 _normalize_token_expiry(token) 

2223 

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) 

2236 

2237 session = await OAuthSessions.create_session( 

2238 user_id=user.id, 

2239 provider=provider, 

2240 token=token, 

2241 db=db, 

2242 ) 

2243 

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 ) 

2253 

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}') 

2259 

2260 return response 

2261 

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 

2270 

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 

2277 

2278 if not logout_token: 

2279 return JSONResponse( 

2280 status_code=400, 

2281 content={'error': 'invalid_request', 'error_description': 'Missing logout_token parameter'}, 

2282 ) 

2283 

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 ) 

2294 

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 ) 

2300 

2301 # 3. Find the configured provider whose issuer matches the token 

2302 matched_provider = None 

2303 matched_client = None 

2304 matched_jwks_uri = None 

2305 

2306 for provider_name in OAUTH_PROVIDERS: 

2307 client = self.get_client(provider_name) 

2308 if not client: 

2309 continue 

2310 

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 

2316 

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 

2322 

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 ) 

2332 

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') 

2338 

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

2343 

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') 

2350 

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 ) 

2373 

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 ) 

2382 

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 ) 

2390 

2391 # 7. Extract sub and/or sid — at least one must be present 

2392 sub = claims.get('sub') 

2393 sid = claims.get('sid') 

2394 

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 ) 

2401 

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) 

2408 

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) 

2411 

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={}) 

2417 

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 ) 

2425 

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) 

2431 

2432 if redis: 

2433 await revoke_user_tokens(request, user.id) 

2434 revoked_count += 1 

2435 

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 ) 

2443 

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={})