This commit is contained in:
yucheng-berri 2026-10-05 12:24:09 -07:00 • committed by GitHub
commit ae4c5bb6ae
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 108 additions and 16 deletions

View file

@ -2320,6 +2320,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(
@ -2333,10 +2337,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] = []
@ -2352,19 +2355,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,
@ -2379,7 +2382,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),
)
@ -2440,6 +2443,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,
@ -2896,6 +2900,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
@ -3524,7 +3539,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(
@ -3595,6 +3610,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,
@ -3647,7 +3664,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"],

View file

@ -1305,6 +1305,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():
"""
@ -3174,7 +3248,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 = {
@ -3344,7 +3418,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)
@ -7402,11 +7476,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",