feat(oauth): preserve redirect_url through OAuth state for post-login… (#38900)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
林玮 (Jade Lin) 2026-07-14 16:23:39 +08:00 committed by GitHub
parent f71a30bebf
commit 82ff93cbdd
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
8 changed files with 179 additions and 7 deletions

View File

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

View File

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

View File

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

View File

@ -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?...")

View File

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

View File

@ -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=...")

View File

@ -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",
}

View File

@ -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"),
[