From a3654d6a7d5b4749e2f418257aefc4a80708ffd9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97=E7=8E=AE=20=28Jade=20Lin=29?= Date: Wed, 22 Jul 2026 13:59:04 +0800 Subject: [PATCH] fix(oauth): reauth after accepting invitation (#39366) --- api/controllers/console/auth/oauth.py | 56 ++++++---- .../controllers/console/auth/test_oauth.py | 13 ++- .../console/auth/test_oauth_redirect.py | 101 +++++++++++++++++- 3 files changed, 146 insertions(+), 24 deletions(-) diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index 2160f3e38ec..a49cf47eaf6 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -6,6 +6,7 @@ from flask import current_app, redirect, request from flask_restx import Resource from pydantic import BaseModel, Field from werkzeug.exceptions import Unauthorized +from werkzeug.wrappers import Response from configs import dify_config from constants.languages import languages @@ -127,6 +128,20 @@ def _preferred_interface_language(language: str | None = None) -> str: return languages[0] +def _redirect_with_console_session(account: Account, target_url: str) -> Response: + """Create a console session and attach its cookies to a redirect response.""" + token_pair = AccountService.login( + account=account, + session=db.session(), + ip_address=extract_remote_ip(request), + ) + response = redirect(target_url) + set_access_token_to_cookie(request, response, token_pair.access_token) + set_refresh_token_to_cookie(request, response, token_pair.refresh_token) + set_csrf_token_to_cookie(request, response, token_pair.csrf_token) + return response + + @console_ns.route("/oauth/login/") class OAuthLogin(Resource): @console_ns.doc("oauth_login") @@ -195,16 +210,26 @@ class OAuthCallback(Resource): return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={urllib.parse.quote(str(e))}") if invite_token and RegisterService.is_valid_invite_token(invite_token): - invitation = RegisterService.get_invitation_by_token(token=invite_token) - if invitation: - invitation_email = invitation.get("email", None) - invitation_email_normalized = ( - invitation_email.lower() if isinstance(invitation_email, str) else invitation_email - ) - if invitation_email_normalized != user_info.email.lower(): - return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Invalid invitation token.") + invitation = RegisterService.get_invitation_if_token_valid( + None, + None, + invite_token, + session=db.session(), + ) + if not invitation: + return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Invalid invitation token.") + if invitation["data"]["email"].lower() != user_info.email.lower(): + message = "This invitation was sent to another account. Please sign in with the invited account." + query = urllib.parse.urlencode({"message": message, "invite_token": invite_token}) + return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?{query}") - return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin/invite-settings?invite_token={invite_token}") + account = invitation["account"] + if account.status == AccountStatus.BANNED: + return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Account is banned.") + + AccountService.link_account_integrate(provider, user_info.id, account, session=db.session()) + target_url = f"{dify_config.CONSOLE_WEB_URL}/signin/invite-settings?invite_token={invite_token}" + return _redirect_with_console_session(account, target_url) try: account, oauth_new_user = _generate_account(provider, user_info, timezone=timezone, language=language) @@ -239,21 +264,10 @@ class OAuthCallback(Resource): "?message=Workspace not found, please contact system admin to invite you to join in a workspace." ) - token_pair = AccountService.login( - account=account, - session=db.session(), - ip_address=extract_remote_ip(request), - ) - 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) - set_refresh_token_to_cookie(request, response, token_pair.refresh_token) - set_csrf_token_to_cookie(request, response, token_pair.csrf_token) - return response + return _redirect_with_console_session(account, target_url) def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) -> Account | None: diff --git a/api/tests/unit_tests/controllers/console/auth/test_oauth.py b/api/tests/unit_tests/controllers/console/auth/test_oauth.py index 6964157189d..a32cac0225f 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_oauth.py +++ b/api/tests/unit_tests/controllers/console/auth/test_oauth.py @@ -250,10 +250,12 @@ class TestOAuthCallback: @patch("controllers.console.auth.oauth.dify_config") @patch("controllers.console.auth.oauth.get_oauth_providers") @patch("controllers.console.auth.oauth.RegisterService") + @patch("controllers.console.auth.oauth.AccountService") @patch("controllers.console.auth.oauth.redirect") def test_invitation_comparison_is_case_insensitive( self, mock_redirect, + mock_account_service, mock_register_service, mock_get_providers, mock_config, @@ -267,13 +269,20 @@ class TestOAuthCallback: ) mock_get_providers.return_value = {"github": oauth_setup["provider"]} mock_register_service.is_valid_invite_token.return_value = True - mock_register_service.get_invitation_by_token.return_value = {"email": "user@example.com"} + mock_register_service.get_invitation_if_token_valid.return_value = { + "account": oauth_setup["account"], + "data": {"email": "user@example.com"}, + "tenant": MagicMock(), + } + mock_account_service.login.return_value = oauth_setup["token_pair"] state = encode_oauth_state(invite_token="invite123", timezone="Asia/Shanghai") with app.test_request_context(f"/auth/oauth/github/callback?code=test_code&state={state}"): resource.get("github") - mock_register_service.get_invitation_by_token.assert_called_once_with(token="invite123") + mock_register_service.get_invitation_if_token_valid.assert_called_once_with( + None, None, "invite123", session=ANY + ) mock_redirect.assert_called_once_with("http://localhost:3000/signin/invite-settings?invite_token=invite123") @pytest.mark.parametrize( 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 index ac1bace882e..e5a4891fe80 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py +++ b/api/tests/unit_tests/controllers/console/auth/test_oauth_redirect.py @@ -1,5 +1,5 @@ import urllib.parse -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -91,3 +91,102 @@ def test_oauth_callback_validates_redirect_url_and_appends_new_user_flag( assert response.headers["Location"] == ( f"{expected_target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}" ) + + +def test_oauth_callback_with_invitation_establishes_console_session(app: Flask) -> 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="Invitee@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(invite_token="invite-token") + + 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.RegisterService") as register_service, + patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account, + patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair) as login, + patch("controllers.console.auth.oauth.TenantService.create_owner_tenant_if_not_exist") as create_workspace, + patch("controllers.console.auth.oauth.set_access_token_to_cookie") as set_access_cookie, + patch("controllers.console.auth.oauth.set_refresh_token_to_cookie") as set_refresh_cookie, + patch("controllers.console.auth.oauth.set_csrf_token_to_cookie") as set_csrf_cookie, + app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"), + ): + register_service.is_valid_invite_token.return_value = True + register_service.get_invitation_if_token_valid.return_value = { + "account": account, + "data": { + "account_id": "account-id", + "email": "invitee@example.com", + "workspace_id": "workspace-id", + }, + "tenant": MagicMock(), + } + + response = OAuthCallback().get("google") + + assert response.status_code == 302 + assert response.headers["Location"] == (f"{CONSOLE_WEB_URL}/signin/invite-settings?invite_token=invite-token") + link_account.assert_called_once_with("google", "google-user-123", account, session=ANY) + login.assert_called_once_with(account=account, session=ANY, ip_address=ANY) + create_workspace.assert_not_called() + set_access_cookie.assert_called_once_with(ANY, response, "dify-access-token") + set_refresh_cookie.assert_called_once_with(ANY, response, "dify-refresh-token") + set_csrf_cookie.assert_called_once_with(ANY, response, "dify-csrf-token") + + +def test_oauth_callback_with_invitation_rejects_another_account(app: Flask) -> 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="another@example.com", + ) + account = MagicMock() + account.status = AccountStatus.ACTIVE + state = encode_oauth_state(invite_token="invite-token") + + 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.RegisterService") as register_service, + patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account, + patch("controllers.console.auth.oauth.AccountService.login") as login, + patch("controllers.console.auth.oauth.set_access_token_to_cookie") as set_access_cookie, + patch("controllers.console.auth.oauth.set_refresh_token_to_cookie") as set_refresh_cookie, + patch("controllers.console.auth.oauth.set_csrf_token_to_cookie") as set_csrf_cookie, + app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"), + ): + register_service.is_valid_invite_token.return_value = True + register_service.get_invitation_if_token_valid.return_value = { + "account": account, + "data": { + "account_id": "account-id", + "email": "invitee@example.com", + "workspace_id": "workspace-id", + }, + "tenant": MagicMock(), + } + + response = OAuthCallback().get("google") + + query = urllib.parse.parse_qs(urllib.parse.urlparse(response.headers["Location"]).query) + assert response.status_code == 302 + assert query["message"] == ["This invitation was sent to another account. Please sign in with the invited account."] + assert query["invite_token"] == ["invite-token"] + link_account.assert_not_called() + login.assert_not_called() + register_service.revoke_token.assert_not_called() + set_access_cookie.assert_not_called() + set_refresh_cookie.assert_not_called() + set_csrf_cookie.assert_not_called()