From f412538756f745b531036840bc0a153e62b0f003 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Thu, 24 Sep 2026 12:10:30 -0400 Subject: [PATCH] refac --- backend/open_webui/routers/auths.py | 12 ++++++++++- backend/open_webui/utils/oauth.py | 33 ++++++++++++----------------- 2 files changed, 25 insertions(+), 20 deletions(-) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index a07d5ed678..0ca2f518f8 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -9,6 +9,7 @@ import urllib import uuid from ssl import CERT_NONE, CERT_REQUIRED, PROTOCOL_TLS +import jwt from aiohttp import BasicAuth, ClientSession from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.responses import JSONResponse, Response @@ -1751,12 +1752,20 @@ async def token_exchange( detail='User not found. Please sign in via the web interface first.', ) + # The provider's userinfo endpoint has already accepted this token. + # Keep an empty dict for opaque tokens so exchange role checks still apply. + token_claims = {} + try: + token_claims = jwt.decode(form_data.token, options={'verify_signature': False}) + except jwt.PyJWTError as e: + log.debug('Token exchange: cannot decode token claims: %s', e) + user = await oauth_manager.update_user_role_from_oauth( request=request, user=user, user_data=user_data, provider=provider, - access_token=form_data.token, + token_claims=token_claims, db=db, ) if await Config.get('oauth.enable_group_mapping'): @@ -1765,6 +1774,7 @@ async def token_exchange( user=user, user_data=user_data, default_permissions=await Config.get('user.permissions'), + token_claims=token_claims, db=db, ) diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 25c6a86a99..6d7bf67d07 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -189,7 +189,7 @@ def _default_value(value): return getattr(value, 'value', value) -def _get_roles_claim(claims: dict, claim: str) -> list | str | int | None: +def _get_claim(claims: dict, claim: str) -> list | str | int | None: """Read nested or flat claims, preserving explicit empty values and zero.""" value = claims for key in claim.split('.'): @@ -1542,7 +1542,7 @@ class OAuthManager: log.error(f'Exception during token refresh for provider {provider}: {e}') return None - async def get_user_role(self, user, user_data, *, access_token: str | None = None): + async def get_user_role(self, user, user_data, *, token_claims: dict | None = None): auth_config = await get_oauth_runtime_config() user_count = await Users.get_num_users() if user and user_count == 1: @@ -1567,14 +1567,9 @@ class OAuthManager: role = user.role if user else auth_config.DEFAULT_USER_ROLE if oauth_claim: - claim_data = _get_roles_claim(user_data, oauth_claim) - if claim_data is None and access_token is not None: - # The exchange endpoint has already validated this token with the provider's userinfo endpoint. - try: - token_claims = jwt.decode(access_token, options={'verify_signature': False}) - claim_data = _get_roles_claim(token_claims, oauth_claim) - except jwt.PyJWTError as e: - log.debug('Token exchange: cannot decode token claims: %s', e) + claim_data = _get_claim(user_data, oauth_claim) + if claim_data is None and token_claims is not None: + claim_data = _get_claim(token_claims, oauth_claim) if isinstance(claim_data, list): oauth_roles = claim_data @@ -1587,7 +1582,7 @@ class OAuthManager: elif isinstance(claim_data, int): oauth_roles = [str(claim_data)] - if access_token is not None and not oauth_roles and oauth_allowed_roles and '*' not in oauth_allowed_roles: + if token_claims is not None and not oauth_roles and oauth_allowed_roles and '*' not in oauth_allowed_roles: log.warning('Token exchange denied: no readable roles claim in userinfo or the token') raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED) @@ -1637,10 +1632,10 @@ class OAuthManager: user_data, provider, *, - access_token: str | None = None, + token_claims: dict | None = None, db=None, ): - determined_role = await self.get_user_role(user, user_data, access_token=access_token) + determined_role = await self.get_user_role(user, user_data, token_claims=token_claims) if user.role == determined_role: return user @@ -1658,7 +1653,9 @@ class OAuthManager: return user - async def update_user_groups(self, request, user, user_data, default_permissions, db=None): + async def update_user_groups( + self, request, user, user_data, default_permissions, db=None, *, token_claims: dict | None = None + ): auth_config = await get_oauth_runtime_config() log.debug('Running OAUTH Group management') oauth_claim = auth_config.OAUTH_GROUPS_CLAIM @@ -1666,12 +1663,10 @@ class OAuthManager: blocked_groups = _parse_blocked_groups(auth_config.OAUTH_BLOCKED_GROUPS) user_oauth_groups = [] - # Nested claim search for groups claim if oauth_claim: - claim_data = user_data - nested_claims = oauth_claim.split('.') - for nested_claim in nested_claims: - claim_data = claim_data.get(nested_claim, {}) + claim_data = _get_claim(user_data, oauth_claim) + if claim_data is None and token_claims is not None: + claim_data = _get_claim(token_claims, oauth_claim) if isinstance(claim_data, list): user_oauth_groups = claim_data