diff --git a/lib/galaxy/authnz/managers.py b/lib/galaxy/authnz/managers.py index f2a2c33825e..e45c91df631 100644 --- a/lib/galaxy/authnz/managers.py +++ b/lib/galaxy/authnz/managers.py @@ -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) diff --git a/lib/galaxy/authnz/xsd/oidc_backends_config.xsd b/lib/galaxy/authnz/xsd/oidc_backends_config.xsd index e62e9804805..d59eca22fa8 100644 --- a/lib/galaxy/authnz/xsd/oidc_backends_config.xsd +++ b/lib/galaxy/authnz/xsd/oidc_backends_config.xsd @@ -65,6 +65,15 @@ + + + + Require the user to refresh their session (via refresh token) + when the access token expires. Users will be required to reauthenticate + if refreshing fails. + + + diff --git a/lib/galaxy/config/sample/oidc_backends_config.xml.sample b/lib/galaxy/config/sample/oidc_backends_config.xml.sample index 0e8ea9cf557..1038b41b050 100644 --- a/lib/galaxy/config/sample/oidc_backends_config.xml.sample +++ b/lib/galaxy/config/sample/oidc_backends_config.xml.sample @@ -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 _______________ diff --git a/lib/galaxy/webapps/base/webapp.py b/lib/galaxy/webapps/base/webapp.py index aea09627b0e..76980af1c34 100644 --- a/lib/galaxy/webapps/base/webapp.py +++ b/lib/galaxy/webapps/base/webapp.py @@ -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 diff --git a/test/unit/authnz/test_authnz.py b/test/unit/authnz/test_authnz.py index 1b72b4756b9..9d17188df0f 100644 --- a/test/unit/authnz/test_authnz.py +++ b/test/unit/authnz/test_authnz.py @@ -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 = """ $galaxy_url/authnz/keycloak/callback {enable_idp_logout} {require_create_confirmation} + {require_session_refresh} {accepted_audiences} {username_key} @@ -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 = """ + + + https://example.com + abcd1234 + abcdef99999 + $galaxy_url/authnz/oidc/callback + + + """ + 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 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