From 60bb33af720312606b687f81815aa6cbe1c54011 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 16:02:12 -0700 Subject: [PATCH] fix(sso): reject empty SSO identities --- litellm/proxy/management_endpoints/ui_sso.py | 35 ++++++-- .../proxy/management_endpoints/test_ui_sso.py | 89 +++++++++++++++++-- 2 files changed, 108 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 19444dfe33b..c3a478e7907 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2282,6 +2282,10 @@ async def _complete_cli_sso_callback_session( ): from fastapi.responses import HTMLResponse + effective_user_id: Final = ( + user_defined_values.get("user_id") if user_defined_values is not None else parsed_openid_result.get("user_id") + ) + _require_sso_user_id(effective_user_id) user_id: Final = parsed_openid_result.get("user_id") user_email: Final = parsed_openid_result.get("user_email") user_info: Final = await get_user_info_from_db( @@ -2295,10 +2299,9 @@ async def _complete_cli_sso_callback_session( ) if user_info is None: raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO") - if not user_info.user_id: - raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO") + resolved_user_id: Final = _require_sso_user_id(user_info.user_id) - await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion) + await retain_sso_identity_assertion_for_ema(user_id=resolved_user_id, assertion=sso_assertion) await warn_if_id_jag_assertion_uncaptured(sso_assertion) teams: list[str] = [] @@ -2314,19 +2317,19 @@ async def _complete_cli_sso_callback_session( from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject await enroll_microsoft_subject( - request.scope.get("litellm_microsoft_interactive_subject"), user_info.user_id, prisma_client + request.scope.get("litellm_microsoft_interactive_subject"), resolved_user_id, prisma_client ) resolved_teams: Final = _cli_sso_session_teams(team_details) attribution_metadata: Final = build_cli_sso_attribution_metadata(result=result) if attribution_metadata: await _persist_cli_sso_user_metadata( prisma_client=prisma_client, - user_id=cast(str, user_info.user_id), + user_id=resolved_user_id, attribution_metadata=attribution_metadata, ) flow["session_data"] = { - "user_id": cast(str, user_info.user_id), + "user_id": resolved_user_id, "user_role": user_info.user_role, "models": user_info.models if hasattr(user_info, "models") else [], "user_email": user_email, @@ -2341,7 +2344,7 @@ async def _complete_cli_sso_callback_session( verbose_proxy_logger.info( "Stored CLI SSO session for user: %s, teams: %s, num_teams: %s", - user_info.user_id, + resolved_user_id, resolved_teams, len(resolved_teams), ) @@ -2402,6 +2405,7 @@ async def cli_sso_callback( result=result_non_none, parsed_openid_result=parsed_openid_result, ) + _require_sso_user_id(user_defined_values.get("user_id") if user_defined_values is not None else None) SSOAuthenticationHandler.verify_user_in_restricted_sso_group( general_settings=general_settings, @@ -2855,6 +2859,17 @@ def _persist_return_to_cookie(response: Response, return_to: str | None, request ) +def _require_sso_user_id(user_id: str | None) -> str: + """Return a nonblank SSO user id or reject the login with a 401""" + if user_id is None or not user_id.strip(): + verbose_proxy_logger.warning("SSO login rejected: the provider response resolved no user id or email") + raise HTTPException( + status_code=401, + detail="SSO login failed: the identity provider did not return a user id or email for this account", + ) + return user_id + + class SSOAuthenticationHandler: """ Handler for SSO Authentication across all SSO providers @@ -3481,7 +3496,7 @@ class SSOAuthenticationHandler: _last_name: Final = getattr(result, "last_name", "") or "" user_id = _first_name + _last_name - if user_email is not None and (user_id is None or len(user_id) == 0): + if user_email is not None and (user_id is None or not user_id.strip()): user_id = user_email return ParsedOpenIDResult( @@ -3552,6 +3567,8 @@ class SSOAuthenticationHandler: budget_duration=internal_user_budget_duration, ) + _require_sso_user_id(user_defined_values.get("user_id") if user_defined_values is not None else None) + # (IF SET) Verify user is in restricted SSO group SSOAuthenticationHandler.verify_user_in_restricted_sso_group( general_settings=general_settings, @@ -3604,7 +3621,7 @@ class SSOAuthenticationHandler: spend=0, team_id="litellm-dashboard", models=user_defined_values["models"], - user_id=user_defined_values["user_id"], + user_id=_require_sso_user_id(user_defined_values["user_id"]), user_email=user_defined_values["user_email"], user_role=user_defined_values["user_role"], max_budget=user_defined_values["max_budget"], diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py index 7db37588cad..d9110467c98 100644 --- a/tests/unit/proxy/management_endpoints/test_ui_sso.py +++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py @@ -1147,6 +1147,80 @@ def test_get_user_email_and_id_extracts_microsoft_role(): assert parsed.get("user_role") == "proxy_admin_viewer" +@pytest.mark.parametrize( + ("result", "expected_user_id"), + [ + pytest.param( + CustomOpenID(id="entra-object-id", email="a@example.com", provider="microsoft", team_ids=[]), + "entra-object-id", + id="id", + ), + pytest.param(CustomOpenID(email="a@example.com", provider="generic", team_ids=[]), "a@example.com", id="email"), + pytest.param( + CustomOpenID(first_name="Ada", last_name="Lovelace", provider="generic", team_ids=[]), + "AdaLovelace", + id="names", + ), + ], +) +def test_get_user_email_and_id_resolves_identity(result, expected_user_id): + parsed = SSOAuthenticationHandler._get_user_email_and_id_from_result(result=result, generic_client_id=None) + + assert parsed.get("user_id") == expected_user_id + + +def _empty_identity_redirect_patches(custom_sso, user_info_from_db, generate_key): + mock_request = MagicMock(spec=Request) + mock_request.scope = {} + mock_request.base_url = "http://localhost:4000/" + mock_request.cookies = {} + stack = ExitStack() + for target, value in ( + ("litellm.proxy.utils.get_prisma_client_or_throw", MagicMock(return_value=MagicMock())), + ("litellm.proxy.proxy_server.master_key", "sk-master"), + ("litellm.proxy.proxy_server.general_settings", {}), + ("litellm.proxy.proxy_server.premium_user", False), + ("litellm.proxy.proxy_server.user_custom_sso", custom_sso), + ("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ("litellm.proxy.proxy_server.redis_usage_cache", None), + ("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ("litellm.proxy.proxy_server.generate_key_helper_fn", generate_key), + ("litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", user_info_from_db), + ): + stack.enter_context(patch(target, value)) # test-quality-ok: endpoint reads proxy globals and DB helpers + return mock_request, stack + + +@pytest.mark.asyncio +async def test_redirect_from_openid_allows_custom_sso_to_resolve_missing_provider_identity(): + async def custom_sso(result): + return { + "models": [], + "user_id": result.extra_fields["employee_id"], + "user_email": None, + "user_role": None, + "max_budget": None, + "budget_duration": None, + } + + user_info_from_db = AsyncMock(side_effect=RuntimeError("stop after identity resolution")) + generate_key = AsyncMock() + mock_request, stack = _empty_identity_redirect_patches(custom_sso, user_info_from_db, generate_key) + result = CustomOpenID(provider="generic", team_ids=[], extra_fields={"employee_id": "mapped-user"}) + + with stack, pytest.raises(RuntimeError, match="stop after identity resolution"): + await SSOAuthenticationHandler.get_redirect_response_from_openid( + result=result, + request=mock_request, + received_response=None, + generic_client_id=None, + ui_access_mode=None, + ) + + assert user_info_from_db.await_args.kwargs["user_defined_values"]["user_id"] == "mapped-user" + generate_key.assert_not_awaited() + + @pytest.mark.asyncio async def test_get_user_info_from_db_user_exists(): """ @@ -3016,7 +3090,7 @@ class TestCLIKeyRegenerationFlow: teams=[], models=[], ) - mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"} + mock_sso_result = CustomOpenID(id="test-user-123", email="test@example.com", team_ids=[]) mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { @@ -3186,7 +3260,7 @@ class TestCLIKeyRegenerationFlow: ) # Mock SSO result - mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"} + mock_sso_result = CustomOpenID(id="test-user-123", email="test@example.com", team_ids=[]) # Mock cache mock_cache = MagicMock(redis_cache=None) @@ -7244,11 +7318,12 @@ class TestCliSsoAttributionMetadata: teams=["team1"], models=["gpt-4"], ) - mock_sso_result = { - "user_email": "test@example.com", - "user_id": "test-user-123", - "employment_type": "contractor", - } + mock_sso_result = CustomOpenID( + id="test-user-123", + email="test@example.com", + team_ids=[], + extra_fields={"employment_type": "contractor"}, + ) mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash",