mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(sso): reject empty SSO identities
This commit is contained in:
parent
4b1d9bf148
commit
60bb33af72
2 changed files with 108 additions and 16 deletions
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue