mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 02:28:30 +08:00
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:
parent
f71a30bebf
commit
82ff93cbdd
@ -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)
|
||||
|
||||
@ -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)}"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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?...")
|
||||
|
||||
|
||||
@ -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()}"
|
||||
)
|
||||
@ -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=...")
|
||||
|
||||
|
||||
@ -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",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@ -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"),
|
||||
[
|
||||
|
||||
Loading…
Reference in New Issue
Block a user