mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix(oauth): reauth after accepting invitation (#39366)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user