From 82ff93cbdd8a9e83d3e59a8626b32299da9a1b9a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97=E7=8E=AE=20=28Jade=20Lin=29?= Date: Tue, 14 Jul 2026 16:23:39 +0800 Subject: [PATCH] =?UTF-8?q?feat(oauth):=20preserve=20redirect=5Furl=20thro?= =?UTF-8?q?ugh=20OAuth=20state=20for=20post-login=E2=80=A6=20(#38900)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/controllers/console/auth/oauth.py | 40 +++++++- api/libs/oauth.py | 21 ++++- api/openapi/markdown/console-openapi.md | 2 + .../controllers/console/auth/test_oauth.py | 3 + .../console/auth/test_oauth_redirect.py | 93 +++++++++++++++++++ .../console/auth/test_oauth_timezone.py | 1 + api/tests/unit_tests/libs/test_oauth_base.py | 10 +- .../unit_tests/libs/test_oauth_clients.py | 16 ++++ 8 files changed, 179 insertions(+), 7 deletions(-) create mode 100644 api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index ffda5b09840..c4c3c5642d5 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -38,6 +38,7 @@ class OAuthLoginQuery(BaseModel): invite_token: str | None = Field(default=None, description="Optional invitation token") timezone: str | None = Field(default=None, description="Preferred timezone") language: str | None = Field(default=None, description="Preferred interface language") + redirect_url: str | None = Field(default=None, description="Relative page to resume after login") class OAuthCallbackQuery(BaseModel): @@ -87,6 +88,36 @@ def _validated_language(value: str | None) -> str | None: return None +def _url_origin(url: str) -> tuple[str, str, int] | None: + parsed_url = urllib.parse.urlsplit(url) + if parsed_url.scheme not in {"http", "https"} or parsed_url.hostname is None: + return None + + try: + port = parsed_url.port + except ValueError: + return None + + if port is None: + port = 443 if parsed_url.scheme == "https" else 80 + return parsed_url.scheme, parsed_url.hostname, port + + +def _get_redirect_target(redirect_url: str | None) -> str: + if not redirect_url: + return dify_config.CONSOLE_WEB_URL + + parsed_url = urllib.parse.urlsplit(redirect_url) + normalized_path = redirect_url.lstrip().replace("\\", "/") + if not parsed_url.scheme and not parsed_url.netloc and not normalized_path.startswith("//"): + return redirect_url + + redirect_origin = _url_origin(redirect_url) + if redirect_origin is not None and redirect_origin == _url_origin(dify_config.CONSOLE_WEB_URL): + return redirect_url + return dify_config.CONSOLE_WEB_URL + + def _preferred_interface_language(language: str | None = None) -> str: if language: return language @@ -109,6 +140,7 @@ class OAuthLogin(Resource): invite_token = request.args.get("invite_token") or None timezone = _validated_timezone(request.args.get("timezone") or None) language = _validated_language(request.args.get("language") or None) + redirect_url = request.args.get("redirect_url") or None OAUTH_PROVIDERS = get_oauth_providers() with current_app.app_context(): oauth_provider = OAUTH_PROVIDERS.get(provider) @@ -119,6 +151,7 @@ class OAuthLogin(Resource): invite_token=invite_token, timezone=timezone, language=language, + redirect_url=redirect_url, ) return redirect(auth_url) @@ -144,6 +177,7 @@ class OAuthCallback(Resource): invite_token = oauth_state.get("invite_token") timezone = _validated_timezone(oauth_state.get("timezone")) language = _validated_language(oauth_state.get("language")) + redirect_url = oauth_state.get("redirect_url") if not code: return {"error": "Authorization code is required"}, 400 @@ -212,9 +246,9 @@ class OAuthCallback(Resource): ip_address=extract_remote_ip(request), ) - base_url = dify_config.CONSOLE_WEB_URL - query_char = "&" if "?" in base_url else "?" - target_url = f"{base_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}" + target_url = _get_redirect_target(redirect_url) + query_char = "&" if "?" in target_url else "?" + target_url = f"{target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}" response = redirect(target_url) set_access_token_to_cookie(request, response, token_pair.access_token) diff --git a/api/libs/oauth.py b/api/libs/oauth.py index 687fd086572..79796e5f4c0 100644 --- a/api/libs/oauth.py +++ b/api/libs/oauth.py @@ -34,6 +34,7 @@ class OAuthState(TypedDict, total=False): invite_token: str timezone: str language: str + redirect_url: str class GitHubEmailRecord(TypedDict, total=False): @@ -72,6 +73,7 @@ def encode_oauth_state( invite_token: str | None = None, timezone: str | None = None, language: str | None = None, + redirect_url: str | None = None, ) -> str | None: state: OAuthState = {} if invite_token: @@ -80,6 +82,8 @@ def encode_oauth_state( state["timezone"] = timezone if language: state["language"] = language + if redirect_url: + state["redirect_url"] = redirect_url if not state: return None @@ -122,6 +126,7 @@ class OAuth: invite_token: str | None = None, timezone: str | None = None, language: str | None = None, + redirect_url: str | None = None, ) -> str: raise NotImplementedError() @@ -151,13 +156,19 @@ class GitHubOAuth(OAuth): invite_token: str | None = None, timezone: str | None = None, language: str | None = None, + redirect_url: str | None = None, ) -> str: params = { "client_id": self.client_id, "redirect_uri": self.redirect_uri, "scope": "user:email", # Request only basic user information } - state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language) + state = encode_oauth_state( + invite_token=invite_token, + timezone=timezone, + language=language, + redirect_url=redirect_url, + ) if state: params["state"] = state return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}" @@ -248,6 +259,7 @@ class GoogleOAuth(OAuth): invite_token: str | None = None, timezone: str | None = None, language: str | None = None, + redirect_url: str | None = None, ) -> str: params = { "client_id": self.client_id, @@ -255,7 +267,12 @@ class GoogleOAuth(OAuth): "redirect_uri": self.redirect_uri, "scope": "openid email", } - state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language) + state = encode_oauth_state( + invite_token=invite_token, + timezone=timezone, + language=language, + redirect_url=redirect_url, + ) if state: params["state"] = state return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}" diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index 7873eebce85..f56b7aa3825 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -7653,6 +7653,7 @@ Initiate OAuth login process | provider | path | OAuth provider name (github/google) | Yes | string | | invite_token | query | Optional invitation token | No | string | | language | query | Preferred interface language | No | string | +| redirect_url | query | Relative page to resume after login | No | string | | timezone | query | Preferred timezone | No | string | #### Responses @@ -19226,6 +19227,7 @@ Coarse node-level status used by Inspector to pick a banner. | ---- | ---- | ----------- | -------- | | invite_token | string | Optional invitation token | No | | language | string | Preferred interface language | No | +| redirect_url | string | Relative page to resume after login | No | | timezone | string | Preferred timezone | No | #### OAuthProviderAccountResponse diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py index 484ca71ca59..ced586a815c 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py @@ -105,6 +105,7 @@ class TestOAuthLogin: invite_token=expected_token, timezone=None, language=None, + redirect_url=None, ) mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...") @@ -127,6 +128,7 @@ class TestOAuthLogin: invite_token=None, timezone="Asia/Shanghai", language=None, + redirect_url=None, ) mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...") @@ -149,6 +151,7 @@ class TestOAuthLogin: invite_token=None, timezone=None, language="zh-Hans", + redirect_url=None, ) mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...") diff --git a/api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py b/api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py new file mode 100644 index 00000000000..ac1bace882e --- /dev/null +++ b/api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py @@ -0,0 +1,93 @@ +import urllib.parse +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask + +from controllers.console.auth.oauth import OAuthCallback, OAuthLogin +from libs.oauth import OAuthUserInfo, encode_oauth_state +from models.account import AccountStatus + +REDIRECT_URL = "/apps?category=workflow" +CONSOLE_WEB_URL = "https://console.example.com" + + +@pytest.fixture +def app() -> Flask: + app = Flask(__name__) + app.config["TESTING"] = True + return app + + +def test_oauth_login_passes_relative_redirect_url_through(app: Flask) -> None: + oauth_provider = MagicMock() + oauth_provider.get_authorization_url.return_value = "https://accounts.google.com/o/oauth2/v2/auth?state=..." + query = urllib.parse.urlencode({"redirect_url": REDIRECT_URL}) + + with ( + patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}), + app.test_request_context(f"/oauth/login/google?{query}"), + ): + response = OAuthLogin().get("google") + + oauth_provider.get_authorization_url.assert_called_once_with( + invite_token=None, + timezone=None, + language=None, + redirect_url=REDIRECT_URL, + ) + assert response.status_code == 302 + assert response.headers["Location"] == "https://accounts.google.com/o/oauth2/v2/auth?state=..." + + +@pytest.mark.parametrize( + ("redirect_url", "expected_target_url"), + [ + (REDIRECT_URL, REDIRECT_URL), + (f"{CONSOLE_WEB_URL}{REDIRECT_URL}", f"{CONSOLE_WEB_URL}{REDIRECT_URL}"), + ("https://console.example.com.malicious.example/apps", CONSOLE_WEB_URL), + ("//malicious.example.com/apps", CONSOLE_WEB_URL), + ("///malicious.example.com/apps", CONSOLE_WEB_URL), + (r"\\malicious.example.com/apps", CONSOLE_WEB_URL), + ], +) +@pytest.mark.parametrize("oauth_new_user", [False, True]) +def test_oauth_callback_validates_redirect_url_and_appends_new_user_flag( + app: Flask, + redirect_url: str, + expected_target_url: str, + oauth_new_user: bool, +) -> None: + oauth_provider = MagicMock() + oauth_provider.get_access_token.return_value = "google-access-token" + oauth_provider.get_user_info.return_value = OAuthUserInfo( + id="google-user-123", + name="Test User", + email="test@example.com", + ) + account = MagicMock() + account.status = AccountStatus.ACTIVE + token_pair = MagicMock() + token_pair.access_token = "dify-access-token" + token_pair.refresh_token = "dify-refresh-token" + token_pair.csrf_token = "dify-csrf-token" + state = encode_oauth_state(redirect_url=redirect_url) + + with ( + patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}), + patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL), + patch("controllers.console.auth.oauth._generate_account", return_value=(account, oauth_new_user)), + patch("controllers.console.auth.oauth.TenantService.create_owner_tenant_if_not_exist"), + patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair), + patch("controllers.console.auth.oauth.set_access_token_to_cookie"), + patch("controllers.console.auth.oauth.set_refresh_token_to_cookie"), + patch("controllers.console.auth.oauth.set_csrf_token_to_cookie"), + app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"), + ): + response = OAuthCallback().get("google") + + assert response.status_code == 302 + query_char = "&" if "?" in expected_target_url else "?" + assert response.headers["Location"] == ( + f"{expected_target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}" + ) diff --git a/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py b/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py index 8b3de6a39e5..75545cc27e7 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py +++ b/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py @@ -33,6 +33,7 @@ def test_oauth_login_passes_language_and_timezone_to_authorization_url( invite_token=None, timezone="Asia/Shanghai", language="zh-Hans", + redirect_url=None, ) mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?state=...") diff --git a/api/tests/unit_tests/libs/test_oauth_base.py b/api/tests/unit_tests/libs/test_oauth_base.py index 1c0066ed9ae..506a54f7ce8 100644 --- a/api/tests/unit_tests/libs/test_oauth_base.py +++ b/api/tests/unit_tests/libs/test_oauth_base.py @@ -19,13 +19,19 @@ def test_oauth_base_methods_raise_not_implemented(): oauth._transform_user_info({}) -def test_oauth_state_round_trips_invite_token_timezone_and_language(): - state = encode_oauth_state(invite_token="invite-123", timezone="Asia/Shanghai", language="zh-Hans") +def test_oauth_state_round_trips_login_context(): + state = encode_oauth_state( + invite_token="invite-123", + timezone="Asia/Shanghai", + language="zh-Hans", + redirect_url="/apps?category=workflow", + ) assert decode_oauth_state(state) == { "invite_token": "invite-123", "timezone": "Asia/Shanghai", "language": "zh-Hans", + "redirect_url": "/apps?category=workflow", } diff --git a/api/tests/unit_tests/libs/test_oauth_clients.py b/api/tests/unit_tests/libs/test_oauth_clients.py index b3ecc5a06de..ebd2c9f895d 100644 --- a/api/tests/unit_tests/libs/test_oauth_clients.py +++ b/api/tests/unit_tests/libs/test_oauth_clients.py @@ -70,6 +70,14 @@ class TestGitHubOAuth(BaseOAuthTest): else: assert "state" not in params + def test_should_preserve_redirect_url_in_state(self, oauth): + redirect_url = "/apps?category=workflow" + + url = oauth.get_authorization_url(redirect_url=redirect_url) + _, params = self.parse_auth_url(url) + + assert decode_oauth_state(params["state"][0]) == {"redirect_url": redirect_url} + @pytest.mark.parametrize( ("response_data", "expected_token", "should_raise"), [ @@ -252,6 +260,14 @@ class TestGoogleOAuth(BaseOAuthTest): else: assert "state" not in params + def test_should_preserve_redirect_url_in_state(self, oauth): + redirect_url = "/apps?category=workflow" + + url = oauth.get_authorization_url(redirect_url=redirect_url) + _, params = self.parse_auth_url(url) + + assert decode_oauth_state(params["state"][0]) == {"redirect_url": redirect_url} + @pytest.mark.parametrize( ("response_data", "expected_token", "should_raise"), [