fix(oauth): reauth after accepting invitation (#39366)

This commit is contained in:
林玮 (Jade Lin) 2026-07-22 13:59:04 +08:00 committed by GitHub
parent 855464e7b4
commit a3654d6a7d
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 146 additions and 24 deletions

View File

@ -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/<provider>")
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:

View File

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

View File

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