diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index 74d1ecc38f7..5165fc3003a 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -56,7 +56,7 @@ from models.account import Account from services.account_service import AccountService, InvitationDetailDict, RegisterService, TenantService from services.billing_service import BillingService from services.entities.auth_entities import LoginFailureReason, LoginPayloadBase -from services.errors.account import AccountRegisterError +from services.errors.account import AccountRegisterError, RefreshTokenAccountNotFoundError, RefreshTokenNotFoundError from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError from services.feature_service import FeatureService @@ -359,18 +359,22 @@ class RefreshTokenApi(Resource): try: new_token_pair = AccountService.refresh_token(refresh_token, session=db.session) + except Unauthorized as exc: + return SimpleResultMessageResponse(result="fail", message=exc.description or "Unauthorized.").model_dump( + mode="json" + ), 401 + except (RefreshTokenNotFoundError, RefreshTokenAccountNotFoundError) as exc: + return SimpleResultMessageResponse(result="fail", message=str(exc)).model_dump(mode="json"), 401 - # Create response with new cookies - # response-contract:ignore cookie-bearing Flask response - response = make_response(SimpleResultResponse(result="success").model_dump(mode="json")) + # Create response with new cookies + # response-contract:ignore cookie-bearing Flask response + response = make_response(SimpleResultResponse(result="success").model_dump(mode="json")) - # Update cookies with new tokens - set_csrf_token_to_cookie(request, response, new_token_pair.csrf_token) - set_access_token_to_cookie(request, response, new_token_pair.access_token) - set_refresh_token_to_cookie(request, response, new_token_pair.refresh_token) - return response - except Exception as e: - return SimpleResultMessageResponse(result="fail", message=str(e)).model_dump(mode="json"), 401 + # Update cookies with new tokens + set_csrf_token_to_cookie(request, response, new_token_pair.csrf_token) + set_access_token_to_cookie(request, response, new_token_pair.access_token) + set_refresh_token_to_cookie(request, response, new_token_pair.refresh_token) + return response def _get_account_with_case_fallback(email: str): diff --git a/api/services/account_service.py b/api/services/account_service.py index 9bfa586c457..1b9fd724a71 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -65,6 +65,8 @@ from services.errors.account import ( LinkAccountIntegrateError, MemberNotInTenantError, NoPermissionError, + RefreshTokenAccountNotFoundError, + RefreshTokenNotFoundError, RoleAlreadyAssignedError, TenantNotFoundError, ) @@ -654,11 +656,11 @@ class AccountService: # Verify the refresh token account_id = redis_client.get(AccountService._get_refresh_token_key(refresh_token)) if not account_id: - raise ValueError("Invalid refresh token") + raise RefreshTokenNotFoundError("Invalid refresh token") account = AccountService.load_user(account_id.decode("utf-8"), session) if not account: - raise ValueError("Invalid account") + raise RefreshTokenAccountNotFoundError("Invalid account") # Generate new access token and refresh token new_access_token = AccountService.get_account_jwt_token(account) diff --git a/api/services/errors/account.py b/api/services/errors/account.py index 4d3d150e072..700c1dd4aaf 100644 --- a/api/services/errors/account.py +++ b/api/services/errors/account.py @@ -17,6 +17,14 @@ class AccountPasswordError(BaseServiceError): pass +class RefreshTokenNotFoundError(BaseServiceError): + pass + + +class RefreshTokenAccountNotFoundError(BaseServiceError): + pass + + class AccountNotLinkTenantError(BaseServiceError): pass diff --git a/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py b/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py index 34fff57b0ad..8effbb96887 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py +++ b/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py @@ -13,8 +13,10 @@ from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask from flask_restx import Api +from werkzeug.exceptions import Unauthorized from controllers.console.auth.login import RefreshTokenApi +from services.errors.account import RefreshTokenAccountNotFoundError, RefreshTokenNotFoundError class TestRefreshTokenApi: @@ -98,18 +100,19 @@ class TestRefreshTokenApi: @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) - def test_refresh_fails_with_invalid_token(self, mock_refresh_token, mock_extract_token, app: Flask): + def test_refresh_returns_unauthorized_for_invalid_refresh_token( + self, mock_refresh_token, mock_extract_token, app: Flask + ): """ - Test token refresh failure with invalid refresh token. + Test token refresh maps invalid refresh tokens to unauthorized responses. Verifies that: - - Exception is caught when token is invalid - - 401 status code is returned - - Error message is included in response + - Invalid refresh token validation failures return 401 + - The failure response preserves the validation message """ # Arrange mock_extract_token.return_value = "invalid_refresh_token" - mock_refresh_token.side_effect = Exception("Invalid refresh token") + mock_refresh_token.side_effect = RefreshTokenNotFoundError("Invalid refresh token") # Act with app.test_request_context("/refresh-token", method="POST"): @@ -119,22 +122,21 @@ class TestRefreshTokenApi: # Assert assert status_code == 401 assert response["result"] == "fail" - assert "Invalid refresh token" in response["message"] + assert response["message"] == "Invalid refresh token" @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) - def test_refresh_fails_with_expired_token(self, mock_refresh_token, mock_extract_token, app: Flask): + def test_refresh_returns_unauthorized_for_invalid_account(self, mock_refresh_token, mock_extract_token, app: Flask): """ - Test token refresh failure with expired refresh token. + Test token refresh maps missing accounts to unauthorized responses. Verifies that: - - Expired tokens are rejected - - 401 status code is returned - - Appropriate error handling + - Invalid account validation failures return 401 + - The failure response preserves the validation message """ # Arrange - mock_extract_token.return_value = "expired_refresh_token" - mock_refresh_token.side_effect = Exception("Refresh token expired") + mock_extract_token.return_value = "refresh_token_for_missing_account" + mock_refresh_token.side_effect = RefreshTokenAccountNotFoundError("Invalid account") # Act with app.test_request_context("/refresh-token", method="POST"): @@ -144,7 +146,71 @@ class TestRefreshTokenApi: # Assert assert status_code == 401 assert response["result"] == "fail" - assert "expired" in response["message"].lower() + assert response["message"] == "Invalid account" + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_returns_unauthorized_for_banned_account(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh maps banned accounts to unauthorized responses. + + Verifies that: + - Authorization failures raised during account loading return 401 + - The failure response preserves the authorization message + """ + # Arrange + mock_extract_token.return_value = "refresh_token_for_banned_account" + mock_refresh_token.side_effect = Unauthorized("Account is banned.") + + # Act + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + response, status_code = refresh_api.post() + + # Assert + assert status_code == 401 + assert response["result"] == "fail" + assert response["message"] == "Account is banned." + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_propagates_non_whitelisted_value_error(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh preserves non-whitelisted ValueError failures. + + Verifies that: + - Only known refresh-token validation errors are mapped to 401 + - Unexpected ValueError instances continue to propagate + """ + # Arrange + mock_extract_token.return_value = "valid_refresh_token" + mock_refresh_token.side_effect = ValueError("unexpected parse failure") + + # Act & Assert + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + with pytest.raises(ValueError, match="unexpected parse failure"): + refresh_api.post() + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_propagates_unexpected_service_errors(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh preserves unexpected service failures. + + Verifies that: + - Operational errors are not misreported as authentication failures + - The original exception is preserved for higher-level error handling + """ + # Arrange + mock_extract_token.return_value = "valid_refresh_token" + mock_refresh_token.side_effect = RuntimeError("redis unavailable") + + # Act & Assert + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + with pytest.raises(RuntimeError, match="redis unavailable"): + refresh_api.post() @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True)