Merge pull request #22449 from AustralianBioCommons/oidc-require-refresh

Require logging in again when OIDC tokens can't be refreshed
This commit is contained in:
Nuwan Goonasekera
2026-05-01 01:01:46 +05:30
committed by GitHub
5 changed files with 365 additions and 18 deletions
+75 -15
View File
@@ -1,11 +1,18 @@
from __future__ import annotations
import builtins
import logging
from typing import TYPE_CHECKING
from typing import (
Optional,
TYPE_CHECKING,
TypedDict,
)
import jwt as pyjwt
from social_core.exceptions import (
AuthAlreadyAssociated,
AuthCanceled,
AuthForbidden,
AuthTokenError,
)
@@ -13,6 +20,7 @@ from galaxy import (
exceptions,
model,
)
from galaxy.model import UserAuthnzToken
from galaxy.util import (
asbool,
etree,
@@ -32,6 +40,7 @@ from .psa_authnz import (
if TYPE_CHECKING:
from galaxy.managers.context import ProvidesAppContext
from galaxy.webapps.base.webapp import GalaxyWebTransaction
OIDC_BACKEND_SCHEMA = resource_path(__name__, "xsd/oidc_backends_config.xsd")
@@ -45,6 +54,11 @@ DEFAULT_OIDC_IDP_ICONS = {
}
class RefreshResult(TypedDict):
refreshed: bool
reauthentication_required: bool
class AuthnzManager:
def __init__(self, app, oidc_config_file, oidc_backends_config_file):
"""
@@ -163,6 +177,8 @@ class AuthnzManager:
rtv["label"] = config_xml.find("label").text
if config_xml.find("require_create_confirmation") is not None:
rtv["require_create_confirmation"] = asbool(config_xml.find("require_create_confirmation").text)
if config_xml.find("require_session_refresh") is not None:
rtv["require_session_refresh"] = asbool(config_xml.find("require_session_refresh").text)
if config_xml.find("prompt") is not None:
rtv["prompt"] = config_xml.find("prompt").text
if config_xml.find("api_url") is not None:
@@ -226,7 +242,7 @@ class AuthnzManager:
# None, if no allowed idp list is set, and a list of EntityIDs if configured (in oidc_backend)
return self.allowed_idps
def _unify_provider_name(self, provider):
def _unify_provider_name(self, provider: str) -> str | None:
if provider.lower() in self.oidc_backends_config:
return provider.lower()
for k, v in BACKENDS_NAME.items():
@@ -234,9 +250,9 @@ class AuthnzManager:
return k.lower()
return None
def _get_authnz_backend(self, provider: str, idphint=None):
def _get_authnz_backend(self, provider: str, idphint: str | None = None) -> tuple[bool, str, PSAAuthnz | None]:
unified_provider_name = self._unify_provider_name(provider)
if unified_provider_name in self.oidc_backends_config:
if unified_provider_name is not None and unified_provider_name in self.oidc_backends_config:
provider = unified_provider_name
identity_provider_class = self._get_identity_provider_factory(self.oidc_backends_implementation[provider])
try:
@@ -281,29 +297,69 @@ class AuthnzManager:
log.warning(msg)
raise exceptions.ItemAccessibilityException(msg)
def refresh_expiring_oidc_tokens_for_provider(self, trans, auth):
def refresh_expiring_oidc_tokens_for_provider(
self, trans: GalaxyWebTransaction, auth: UserAuthnzToken
) -> RefreshResult:
"""
Refresh expiring OIDC tokens for a specific provider.
Returns:
RefreshResult: A dictionary containing a boolean indicating success, and a boolean
indicating if reauthentication is required
"""
try:
if auth.provider is None:
raise exceptions.AuthenticationFailed("Provider is not set")
success, message, backend = self._get_authnz_backend(auth.provider)
if backend is None:
msg = f"Provider `{auth.provider}` not found"
log.error(msg)
return {"refreshed": False, "reauthentication_required": False}
if success is False:
msg = f"An error occurred when refreshing user token on `{auth.provider}` identity provider: {message}"
log.error(msg)
return False
return {"refreshed": False, "reauthentication_required": False}
refreshed = backend.refresh(trans, auth)
if refreshed:
log.debug(f"Refreshed user token via `{auth.provider}` identity provider")
return True
except Exception:
log.exception("An error occurred when refreshing user token")
return False
return {"refreshed": refreshed, "reauthentication_required": False}
except (AuthTokenError, AuthCanceled, AuthForbidden):
log.warning("Authentication session has expired or is invalid, reauth required.")
return {"refreshed": False, "reauthentication_required": True}
except Exception as e:
log.warning(f"An error occurred when refreshing user token: {e}")
return {"refreshed": False, "reauthentication_required": False}
def refresh_expiring_oidc_tokens(self, trans, user=None):
def refresh_expiring_oidc_tokens(
self, trans: GalaxyWebTransaction, user: Optional[model.User] = None
) -> str | None:
"""
Refresh expiring OIDC tokens for all providers associated with a user.
Returns:
str | None: The provider name if refresh fails and require_session_refresh is enabled, otherwise None
"""
user = trans.user or user
if not isinstance(user, model.User):
return
return None
for auth in user.social_auth or []:
self.refresh_expiring_oidc_tokens_for_provider(trans, auth)
result = self.refresh_expiring_oidc_tokens_for_provider(trans, auth)
if auth.provider is None:
continue
provider = self._unify_provider_name(auth.provider)
if provider is None:
continue
config = self.oidc_backends_config.get(provider, None)
if config is None:
continue
# Redirect to OIDC login if refresh fails and require_session_refresh is enabled
if config.get("require_session_refresh") and result["reauthentication_required"]:
return provider
return None
def authenticate(self, provider, trans, idphint=None):
def authenticate(
self, provider: str, trans: GalaxyWebTransaction, idphint: str | None = None
) -> tuple[bool, str, str | None]:
"""
:type provider: string
:param provider: set the name of the identity provider to be
@@ -314,6 +370,8 @@ class AuthnzManager:
"""
try:
success, message, backend = self._get_authnz_backend(provider, idphint=idphint)
if backend is None:
return False, f"Provider `{provider}` not found", None
if success is False:
return False, message, None
# Check allowed IDPs for providers that support idphint (keycloak, cilogon)
@@ -365,9 +423,11 @@ class AuthnzManager:
log.exception(msg)
return False, msg, (None, None)
def create_user(self, provider: str, token: str, trans: "ProvidesAppContext", login_redirect_url: str):
def create_user(self, provider: str, token: str, trans: ProvidesAppContext, login_redirect_url: str):
try:
success, message, backend = self._get_authnz_backend(provider)
if backend is None:
raise ValueError(f"Provider `{provider}` not found")
if success is False:
return False, message, (None, None)
return success, message, backend.create_user(token, trans, login_redirect_url)
@@ -65,6 +65,15 @@
</xs:documentation>
</xs:annotation>
</xs:element>
<xs:element name="require_session_refresh" minOccurs="0" type="xs:boolean">
<xs:annotation>
<xs:documentation>
Require the user to refresh their session (via refresh token)
when the access token expires. Users will be required to reauthenticate
if refreshing fails.
</xs:documentation>
</xs:annotation>
</xs:element>
<xs:element name="ca_bundle" minOccurs="0" type="xs:string">
<xs:annotation>
<xs:documentation>
@@ -28,6 +28,9 @@ _______________
- require_create_confirmation: A boolean value that decides whether a NewUserConfirmation page shows up.
- require_session_refresh: A boolean value that decides whether failed token refresh requires the
user to reauthenticate with this provider.
IMPORTANT NOTES
_______________
+21 -1
View File
@@ -358,7 +358,12 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo
self._ensure_valid_session(session_cookie)
if hasattr(self.app, "authnz_manager") and self.app.authnz_manager:
self.app.authnz_manager.refresh_expiring_oidc_tokens(self)
# Check for expiring tokens and refresh them. If configured (at the individual provider
# level), require a reauthentication on failed refresh.
reauth_provider = self.app.authnz_manager.refresh_expiring_oidc_tokens(self)
if reauth_provider:
self.handle_user_reauthentication(reauth_provider)
return
if self.galaxy_session:
# When we've authenticated by session, we have to check the
@@ -893,6 +898,21 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo
elif self.webapp.name == "tool_shed":
self.__update_session_cookie(name="galaxycommunitysession")
def handle_user_reauthentication(self, reauth_provider: str) -> None:
"""
Handle user being required to log in again after failed OIDC refresh
"""
log.info("OIDC refresh failed terminally for provider `%s`, forcing re-login", reauth_provider)
if self.galaxy_session:
self.handle_user_logout()
if self.environ.get("is_api_request", False):
self.response.status = 401
self.error_message = "Authentication session expired. Please log in again."
self.user = None
self.galaxy_session = None
else:
self.response.send_redirect(url_for(f"/authnz/{reauth_provider}/login", redirect="true", next="/"))
def get_galaxy_session(self):
"""
Return the current galaxy session
+257 -2
View File
@@ -1,12 +1,36 @@
import tempfile
from typing import Optional
from unittest.mock import MagicMock
from typing import (
Any,
cast,
Optional,
)
from unittest.mock import (
MagicMock,
patch,
)
from urllib.parse import urlencode
import pytest
from social_core.backends.base import BaseAuth
from social_core.exceptions import (
AuthCanceled,
AuthForbidden,
AuthTokenError,
)
from social_core.utils import setting_name
from webob.exc import HTTPFound
from galaxy import model
from galaxy.app_unittest_utils import galaxy_mock
from galaxy.authnz.managers import AuthnzManager
from galaxy.util import asbool
from galaxy.web.framework import base as web_framework_base
from galaxy.web.framework.base import Response
from galaxy.webapps.base.webapp import WebApplication
from ..webapps.test_webapp_base import (
CORSParsingMockConfig,
StubGalaxyWebTransaction,
)
@pytest.fixture
@@ -23,6 +47,7 @@ OIDC_BACKEND_CONFIG_TEMPLATE = """<?xml version="1.0"?>
<redirect_uri>$galaxy_url/authnz/keycloak/callback</redirect_uri>
<enable_idp_logout>{enable_idp_logout}</enable_idp_logout>
<require_create_confirmation>{require_create_confirmation}</require_create_confirmation>
<require_session_refresh>{require_session_refresh}</require_session_refresh>
<accepted_audiences>{accepted_audiences}</accepted_audiences>
<username_key>{username_key}</username_key>
</provider>
@@ -52,6 +77,7 @@ def create_backend_config(
client_secret="client_secret",
enable_idp_logout="true",
require_create_confirmation="false",
require_session_refresh="false",
accepted_audiences="https://audience.example.com",
username_key="custom_username",
) -> tuple[str, str]:
@@ -62,6 +88,7 @@ def create_backend_config(
client_secret=client_secret,
enable_idp_logout=enable_idp_logout,
require_create_confirmation=require_create_confirmation,
require_session_refresh=require_session_refresh,
accepted_audiences=accepted_audiences,
username_key=username_key,
)
@@ -70,6 +97,62 @@ def create_backend_config(
return contents, file.name
class FakeRefreshBackend:
refresh_result: bool = False
refresh_exception: Exception | None = None
def __init__(self, provider, oidc_config, oidc_backend_config, app_config):
self.provider = provider
self.oidc_config = oidc_config
self.oidc_backend_config = oidc_backend_config
self.app_config = app_config
def refresh(self, trans, auth):
if self.refresh_exception is not None:
raise self.refresh_exception
return self.refresh_result
FAKE_SOCIAL_AUTH_BACKEND = cast(BaseAuth, MagicMock())
class AuthenticatedStubGalaxyWebTransaction(StubGalaxyWebTransaction):
auth_user: model.User | None = None
def _ensure_valid_session(self, session_cookie: str, create: bool = True) -> None:
self.user = self.auth_user
self.galaxy_session = None
def _authenticate_api(self, session_cookie: str) -> Optional[str]:
self.user = self.auth_user
self.galaxy_session = None
return None
def _make_user_with_social_auth(provider: str = "oidc") -> model.User:
user = model.User(email="user@example.com", password="password")
auth = model.UserAuthnzToken(provider=provider, uid="user-1", extra_data={"refresh_token": "refresh"}, user=user)
user.social_auth.append(auth)
return user
def _make_authnz_manager(
app: Any, provider_name: str = "oidc", require_session_refresh: str = "false"
) -> AuthnzManager:
_, oidc_path = create_oidc_config()
_, backend_path = create_backend_config(
provider_name=provider_name, require_session_refresh=require_session_refresh
)
app.config.oidc = {}
return AuthnzManager(app=app, oidc_config_file=oidc_path, oidc_backends_config_file=backend_path)
def _make_mock_trans_with_user(user: model.User) -> galaxy_mock.MockTrans:
app = galaxy_mock.MockApp()
trans = galaxy_mock.MockTrans(app=app, user=user)
return trans
def test_parse_backend_config(mock_app):
config_values = {
"url": "https://example.com",
@@ -77,6 +160,7 @@ def test_parse_backend_config(mock_app):
"client_secret": "abcd1234",
"enable_idp_logout": "true",
"require_create_confirmation": "false",
"require_session_refresh": "true",
"accepted_audiences": "https://audience.example.com",
"username_key": "custom_username",
}
@@ -93,6 +177,33 @@ def test_parse_backend_config(mock_app):
# Boolean values should be parsed into bools
assert parsed["enable_idp_logout"] == asbool(config_values["enable_idp_logout"])
assert parsed["require_create_confirmation"] == asbool(config_values["require_create_confirmation"])
assert parsed["require_session_refresh"] == asbool(config_values["require_session_refresh"])
def test_parse_backend_config_bool_defaults(mock_app):
# XML config without boolean fields
config = """<?xml version="1.0"?>
<OIDC>
<provider name="oidc">
<url>https://example.com</url>
<client_id>abcd1234</client_id>
<client_secret>abcdef99999</client_secret>
<redirect_uri>$galaxy_url/authnz/oidc/callback</redirect_uri>
</provider>
</OIDC>
"""
config_file = tempfile.NamedTemporaryFile(mode="w", delete=False)
config_file.write(config)
config_file.flush()
config_file.close()
oidc_contents, oidc_path = create_oidc_config()
manager = AuthnzManager(app=mock_app, oidc_config_file=oidc_path, oidc_backends_config_file=config_file.name)
assert isinstance(manager.oidc_backends_config["oidc"], dict)
parsed = manager.oidc_backends_config["oidc"]
# Boolean values should be False by default
assert parsed["enable_idp_logout"] is False
assert parsed.get("require_create_confirmation", False) is False
assert parsed.get("require_session_refresh", False) is False
def test_psa_authnz_config(mock_app):
@@ -199,3 +310,147 @@ def test_missing_idphint_is_none(mock_app):
app_config=mock_app.config,
)
assert psa.config.get("IDPHINT") is None, "IDPHINT must be None when <idphint> is absent from XML"
def test_refresh_expiring_oidc_tokens_returns_none_after_successful_refresh(mock_app):
user = _make_user_with_social_auth()
trans = _make_mock_trans_with_user(user)
manager = _make_authnz_manager(trans.app)
FakeRefreshBackend.refresh_result = True
FakeRefreshBackend.refresh_exception = None
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
reauth_provider = manager.refresh_expiring_oidc_tokens(cast(Any, trans))
assert reauth_provider is None
@pytest.mark.parametrize(
"refresh_exception",
[
AuthTokenError(backend=FAKE_SOCIAL_AUTH_BACKEND),
AuthCanceled(backend=FAKE_SOCIAL_AUTH_BACKEND),
AuthForbidden(backend=FAKE_SOCIAL_AUTH_BACKEND),
],
)
def test_refresh_expiring_oidc_tokens_returns_provider_on_required_terminal_refresh_failure(
mock_app, refresh_exception
):
user = _make_user_with_social_auth()
trans = _make_mock_trans_with_user(user)
manager = _make_authnz_manager(trans.app, require_session_refresh="true")
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = refresh_exception
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
reauth_provider = manager.refresh_expiring_oidc_tokens(cast(Any, trans))
assert reauth_provider == "oidc"
def test_refresh_expiring_oidc_tokens_returns_none_on_optional_terminal_refresh_failure(mock_app):
user = _make_user_with_social_auth()
trans = _make_mock_trans_with_user(user)
manager = _make_authnz_manager(trans.app, require_session_refresh="false")
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = AuthTokenError(backend=FAKE_SOCIAL_AUTH_BACKEND)
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
reauth_provider = manager.refresh_expiring_oidc_tokens(cast(Any, trans))
assert reauth_provider is None
def test_refresh_expiring_oidc_tokens_uses_unified_provider_for_refresh_config(mock_app):
user = _make_user_with_social_auth(provider="google-openidconnect")
trans = _make_mock_trans_with_user(user)
manager = _make_authnz_manager(trans.app, provider_name="google", require_session_refresh="true")
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = AuthTokenError(backend=FAKE_SOCIAL_AUTH_BACKEND)
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
reauth_provider = manager.refresh_expiring_oidc_tokens(cast(Any, trans))
assert reauth_provider == "google"
def test_refresh_expiring_oidc_tokens_returns_none_on_unexpected_refresh_failure(mock_app):
user = _make_user_with_social_auth()
trans = _make_mock_trans_with_user(user)
manager = _make_authnz_manager(trans.app, require_session_refresh="true")
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = RuntimeError("unexpected refresh failure")
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
reauth_provider = manager.refresh_expiring_oidc_tokens(cast(Any, trans))
assert reauth_provider is None
def test_redirects_to_oidc_login_on_terminal_refresh_failure() -> None:
app = cast(Any, galaxy_mock.MockApp())
app.config = CORSParsingMockConfig()
app.authnz_manager = _make_authnz_manager(app, require_session_refresh="true")
webapp = cast(WebApplication, galaxy_mock.MockWebapp(app.security))
environ = galaxy_mock.buildMockEnviron()
AuthenticatedStubGalaxyWebTransaction.auth_user = _make_user_with_social_auth()
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = AuthTokenError(backend=FAKE_SOCIAL_AUTH_BACKEND)
original_send_redirect = Response.send_redirect
redirect_urls: list[str] = []
def capture_send_redirect(self, url: str):
redirect_urls.append(url)
return original_send_redirect(self, url)
def build_test_url(path: str, **query_params):
return f"{path}?{urlencode(query_params)}"
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
with patch.object(web_framework_base.routes, "url_for", side_effect=build_test_url):
with patch.object(Response, "send_redirect", autospec=True, side_effect=capture_send_redirect):
with pytest.raises(HTTPFound):
AuthenticatedStubGalaxyWebTransaction(environ, app, webapp, "session_cookie")
assert redirect_urls
assert "/authnz/oidc/login" in redirect_urls[0]
def test_returns_401_for_api_request_on_terminal_refresh_failure() -> None:
app = cast(Any, galaxy_mock.MockApp())
app.config = CORSParsingMockConfig()
app.authnz_manager = _make_authnz_manager(app, require_session_refresh="true")
webapp = cast(WebApplication, galaxy_mock.MockWebapp(app.security))
environ = galaxy_mock.buildMockEnviron(PATH_INFO="/api/users/current", is_api_request=True)
AuthenticatedStubGalaxyWebTransaction.auth_user = _make_user_with_social_auth()
FakeRefreshBackend.refresh_result = False
FakeRefreshBackend.refresh_exception = AuthTokenError(backend=FAKE_SOCIAL_AUTH_BACKEND)
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
trans = AuthenticatedStubGalaxyWebTransaction(environ, app, webapp, "session_cookie")
assert trans.response.status == 401
assert trans.error_message == "Authentication session expired. Please log in again."
assert trans.user is None
assert trans.galaxy_session is None
def test_allows_api_request_on_successful_oidc_refresh() -> None:
"""
Test that API requests proceed when the refresh succeeds.
"""
app = cast(Any, galaxy_mock.MockApp())
app.config = CORSParsingMockConfig()
app.authnz_manager = _make_authnz_manager(app, require_session_refresh="true")
webapp = cast(WebApplication, galaxy_mock.MockWebapp(app.security))
environ = galaxy_mock.buildMockEnviron(PATH_INFO="/api/users/current", is_api_request=True)
AuthenticatedStubGalaxyWebTransaction.auth_user = _make_user_with_social_auth()
FakeRefreshBackend.refresh_result = True
FakeRefreshBackend.refresh_exception = None
with patch.object(AuthnzManager, "_get_identity_provider_factory", return_value=FakeRefreshBackend):
trans = AuthenticatedStubGalaxyWebTransaction(environ, app, webapp, "session_cookie")
assert trans.response.status == 200
assert trans.error_message is None
assert trans.galaxy_session is None