litellm/tests/unit/proxy/management_endpoints/test_ui_sso.py
devin-ai-integration[bot] 9351463755
fix(proxy): persist SSO display name as user_alias on login (#44065)
* fix(proxy): persist SSO display name as user_alias on login

Generic/Microsoft SSO already parsed the IdP display_name, first_name and last_name into the SSO result, but the user upsert only wrote user_email and user_role, so the Users table never showed a name. Store the display name (first + last as fallback) as user_alias on first login and on later logins of users whose alias is still empty; never overwrite an alias already set. Whitespace-only names are treated as missing.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep stored user_email when SSO login carries no email claim

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* revert: keep stored user_email change, login writes the IdP email as before

Reverts 47c65cecff. A stored email staying eligible for email-based account linking after the IdP stops sending it is not wanted; the PR goes back to the user_alias fix only

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-10-01 17:14:31 -07:00

9661 lines
369 KiB
Python

import asyncio
import json
import logging
import os
from contextlib import ExitStack, asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from fastapi import HTTPException, Request
import litellm
from litellm._uuid import uuid
from litellm.proxy._types import LiteLLM_UserTable, NewUserResponse
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import (
GoogleSSOHandler,
MicrosoftSSOHandler,
SSOAuthenticationHandler,
_setup_team_mappings,
_sync_user_role_from_jwt_role_map,
normalize_email,
process_sso_jwt_access_token,
)
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
MicrosoftGraphAPIUserGroupDirectoryObject,
MicrosoftGraphAPIUserGroupResponse,
MicrosoftServicePrincipalTeam,
TeamMappings,
)
_SSO_PROVIDER_ENV_VARS = (
"DISABLE_ADMIN_UI",
"MICROSOFT_CLIENT_ID",
"GOOGLE_CLIENT_ID",
"GENERIC_CLIENT_ID",
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
)
def _wire_team_create_tx(prisma_client):
"""`/team/new` inserts the team and mirrors it onto the access groups in one transaction,
so a mocked client has to hand its team table back out of `db.tx()`."""
@asynccontextmanager
async def _tx():
yield SimpleNamespace(
litellm_teamtable=prisma_client.db.litellm_teamtable,
query_raw=AsyncMock(return_value=[]),
)
prisma_client.db.tx = lambda *_args, **_kwargs: _tx()
def test_microsoft_sso_handler_openid_from_response_user_principal_name():
# Arrange
# Create a mock response similar to what Microsoft SSO would return
mock_response = {
"userPrincipalName": "test@example.com",
"displayName": "Test User",
"id": "user123",
"givenName": "Test",
"surname": "User",
"some_other_field": "value",
}
expected_team_ids = ["team1", "team2"]
# Act
# Call the method being tested
result = MicrosoftSSOHandler.openid_from_response(
response=mock_response, team_ids=expected_team_ids, user_role=None
)
# Assert
# Check that the result is a CustomOpenID object with the expected values
assert isinstance(result, CustomOpenID)
assert result.email == "test@example.com"
assert result.display_name == "Test User"
assert result.provider == "microsoft"
assert result.id == "user123"
assert result.first_name == "Test"
assert result.last_name == "User"
assert result.team_ids == expected_team_ids
def test_microsoft_sso_handler_openid_from_response():
# Arrange
# Create a mock response similar to what Microsoft SSO would return
mock_response = {
"mail": "test@example.com",
"displayName": "Test User",
"id": "user123",
"givenName": "Test",
"surname": "User",
"some_other_field": "value",
}
expected_team_ids = ["team1", "team2"]
# Act
# Call the method being tested
result = MicrosoftSSOHandler.openid_from_response(
response=mock_response, team_ids=expected_team_ids, user_role=None
)
# Assert
# Check that the result is a CustomOpenID object with the expected values
assert isinstance(result, CustomOpenID)
assert result.email == "test@example.com"
assert result.display_name == "Test User"
assert result.provider == "microsoft"
assert result.id == "user123"
assert result.first_name == "Test"
assert result.last_name == "User"
assert result.team_ids == expected_team_ids
def test_microsoft_sso_handler_with_empty_response():
# Arrange
# Test with None response
# Act
result = MicrosoftSSOHandler.openid_from_response(
response=None, team_ids=[], user_role=None
)
# Assert
assert isinstance(result, CustomOpenID)
assert result.email is None
assert result.display_name is None
assert result.provider == "microsoft"
assert result.id is None
assert result.first_name is None
assert result.last_name is None
assert result.team_ids == []
def test_microsoft_sso_handler_openid_from_response_with_custom_attributes():
"""
Test that MicrosoftSSOHandler.openid_from_response uses custom attribute names
from constants when environment variables are set.
"""
# Arrange
mock_response = {
"custom_email_field": "custom@example.com",
"custom_display_name": "Custom Display Name",
"custom_id_field": "custom_user_123",
"custom_first_name": "CustomFirst",
"custom_last_name": "CustomLast",
}
expected_team_ids = ["team1"]
# Act
with (
patch("litellm.constants.MICROSOFT_USER_EMAIL_ATTRIBUTE", "custom_email_field"),
patch(
"litellm.constants.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE",
"custom_display_name",
),
patch("litellm.constants.MICROSOFT_USER_ID_ATTRIBUTE", "custom_id_field"),
patch(
"litellm.constants.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE", "custom_first_name"
),
patch(
"litellm.constants.MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "custom_last_name"
),
patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_EMAIL_ATTRIBUTE",
"custom_email_field",
),
patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE",
"custom_display_name",
),
patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_ID_ATTRIBUTE",
"custom_id_field",
),
patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_FIRST_NAME_ATTRIBUTE",
"custom_first_name",
),
patch(
"litellm.proxy.management_endpoints.ui_sso.MICROSOFT_USER_LAST_NAME_ATTRIBUTE",
"custom_last_name",
),
):
result = MicrosoftSSOHandler.openid_from_response(
response=mock_response, team_ids=expected_team_ids, user_role=None
)
# Assert
assert isinstance(result, CustomOpenID)
assert result.email == "custom@example.com"
assert result.display_name == "Custom Display Name"
assert result.provider == "microsoft"
assert result.id == "custom_user_123"
assert result.first_name == "CustomFirst"
assert result.last_name == "CustomLast"
assert result.team_ids == expected_team_ids
@pytest.fixture
def stubbed_graph_api(httpx_transport):
with respx.mock:
respx.get(url__regex=r".*graph\.microsoft\.com.*").mock(
return_value=httpx.Response(200, json={"value": []})
)
yield
def test_get_microsoft_callback_response(stubbed_graph_api):
# Arrange
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_response = {
"mail": "microsoft_user@example.com",
"displayName": "Microsoft User",
"id": "msft123",
"givenName": "Microsoft",
"surname": "User",
}
with patch.dict(
os.environ,
{"MICROSOFT_CLIENT_SECRET": "mock_secret", "MICROSOFT_TENANT": "mock_tenant"},
):
mock_verify = AsyncMock(return_value=mock_response)
with patch(
"fastapi_sso.sso.microsoft.MicrosoftSSO.verify_and_process",
new=mock_verify,
):
# Act
result = asyncio.run(
MicrosoftSSOHandler.get_microsoft_callback_response(
request=mock_request,
microsoft_client_id="mock_client_id",
redirect_url="http://mock_redirect_url",
)
)
# Assert
assert isinstance(result, CustomOpenID)
assert result.email == "microsoft_user@example.com"
assert result.display_name == "Microsoft User"
assert result.provider == "microsoft"
assert result.id == "msft123"
assert result.first_name == "Microsoft"
assert result.last_name == "User"
def test_get_microsoft_callback_response_raw_sso_response(stubbed_graph_api):
# Arrange
mock_request = MagicMock(spec=Request)
mock_response = {
"mail": "microsoft_user@example.com",
"displayName": "Microsoft User",
"id": "msft123",
"givenName": "Microsoft",
"surname": "User",
}
with patch.dict(
os.environ,
{"MICROSOFT_CLIENT_SECRET": "mock_secret", "MICROSOFT_TENANT": "mock_tenant"},
):
mock_verify = AsyncMock(return_value=mock_response)
with patch(
"fastapi_sso.sso.microsoft.MicrosoftSSO.verify_and_process",
new=mock_verify,
):
# Act
result = asyncio.run(
MicrosoftSSOHandler.get_microsoft_callback_response(
request=mock_request,
microsoft_client_id="mock_client_id",
redirect_url="http://mock_redirect_url",
return_raw_sso_response=True,
)
)
# Assert
assert isinstance(result, dict)
assert result["mail"] == "microsoft_user@example.com"
assert result["displayName"] == "Microsoft User"
assert result["id"] == "msft123"
assert result["givenName"] == "Microsoft"
assert result["surname"] == "User"
def test_get_google_callback_response():
# Arrange
mock_request = MagicMock(spec=Request)
mock_response = {
"email": "google_user@example.com",
"name": "Google User",
"sub": "google123",
"given_name": "Google",
"family_name": "User",
}
with patch.dict(os.environ, {"GOOGLE_CLIENT_SECRET": "mock_secret"}):
mock_verify = AsyncMock(return_value=mock_response)
with patch(
"fastapi_sso.sso.google.GoogleSSO.verify_and_process", new=mock_verify
):
# Act
result = asyncio.run(
GoogleSSOHandler.get_google_callback_response(
request=mock_request,
google_client_id="mock_client_id",
redirect_url="http://mock_redirect_url",
)
)
# Assert
assert isinstance(result, dict)
assert result.get("email") == "google_user@example.com"
assert result.get("name") == "Google User"
assert result.get("sub") == "google123"
assert result.get("given_name") == "Google"
assert result.get("family_name") == "User"
@pytest.mark.asyncio
async def test_get_user_groups_from_graph_api():
# Arrange
mock_response = {
"@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects",
"value": [
{
"@odata.type": "#microsoft.graph.group",
"id": "group1",
"displayName": "Group 1",
},
{
"@odata.type": "#microsoft.graph.group",
"id": "group2",
"displayName": "Group 2",
},
],
}
async def mock_get(*args, **kwargs):
mock = MagicMock()
mock.json.return_value = mock_response
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
# Act
result = await MicrosoftSSOHandler.get_user_groups_from_graph_api(
access_token="mock_token"
)
# Assert
assert isinstance(result, list)
assert len(result) == 2
assert "group1" in result
assert "group2" in result
@pytest.mark.asyncio
async def test_get_user_groups_empty_response():
# Arrange
mock_response = {
"@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects",
"value": [],
}
async def mock_get(*args, **kwargs):
mock = MagicMock()
mock.json.return_value = mock_response
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
# Act
result = await MicrosoftSSOHandler.get_user_groups_from_graph_api(
access_token="mock_token"
)
# Assert
assert isinstance(result, list)
assert len(result) == 0
@pytest.mark.asyncio
async def test_get_user_groups_error_handling():
# Arrange
async def mock_get(*args, **kwargs):
raise Exception("API Error")
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
# Act
result = await MicrosoftSSOHandler.get_user_groups_from_graph_api(
access_token="mock_token"
)
# Assert
assert isinstance(result, list)
assert len(result) == 0
@pytest.mark.asyncio
async def test_get_user_groups_uses_default_graph_endpoint(monkeypatch):
monkeypatch.delenv("MICROSOFT_GRAPH_ENDPOINT", raising=False)
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token")
assert requested_urls == ["https://graph.microsoft.com/v1.0/me/memberOf"]
@pytest.mark.asyncio
async def test_get_user_groups_uses_configured_graph_endpoint(monkeypatch):
monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0")
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token")
assert requested_urls == ["https://graph.microsoft.us/v1.0/me/memberOf"]
@pytest.mark.asyncio
async def test_get_group_ids_from_service_principal_uses_configured_graph_endpoint(monkeypatch):
monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0")
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
async_client = MagicMock()
async_client.get = mock_get
await MicrosoftSSOHandler.get_group_ids_from_service_principal(
service_principal_id="sp-123",
async_client=async_client,
access_token="mock_token",
)
assert requested_urls == [
"https://graph.microsoft.us/v1.0/servicePrincipals/sp-123/appRoleAssignedTo"
]
@pytest.mark.asyncio
async def test_get_group_ids_from_service_principal_paginates_through_all_pages():
# Arrange
page_one = {
"@odata.nextLink": "https://graph.microsoft.com/v1.0/servicePrincipals/sp-123/appRoleAssignedTo?$skiptoken=page2",
"value": [
{
"principalType": "Group",
"principalId": "group-on-page-1",
"principalDisplayName": "Group On Page 1",
}
],
}
page_two = {
"value": [
{
"principalType": "Group",
"principalId": "group-on-page-2",
"principalDisplayName": "Group On Page 2",
}
],
}
responses = [page_one, page_two]
async def mock_get(url, *args, **kwargs):
mock = MagicMock()
mock.json.return_value = responses.pop(0)
return mock
async_client = MagicMock()
async_client.get = mock_get
# Act
group_ids, teams = await MicrosoftSSOHandler.get_group_ids_from_service_principal(
service_principal_id="sp-123",
async_client=async_client,
access_token="mock_token",
)
# Assert
assert group_ids == ["group-on-page-1", "group-on-page-2"]
assert [team["principalId"] for team in teams] == [
"group-on-page-1",
"group-on-page-2",
]
def test_get_group_ids_from_graph_api_response():
# Arrange
mock_response = MicrosoftGraphAPIUserGroupResponse(
odata_context="https://graph.microsoft.com/v1.0/$metadata#directoryObjects",
odata_nextLink=None,
value=[
MicrosoftGraphAPIUserGroupDirectoryObject(
odata_type="#microsoft.graph.group",
id="group1",
displayName="Group 1",
description=None,
deletedDateTime=None,
roleTemplateId=None,
),
MicrosoftGraphAPIUserGroupDirectoryObject(
odata_type="#microsoft.graph.group",
id="group2",
displayName="Group 2",
description=None,
deletedDateTime=None,
roleTemplateId=None,
),
MicrosoftGraphAPIUserGroupDirectoryObject(
odata_type="#microsoft.graph.group",
id=None, # Test handling of None id
displayName="Invalid Group",
description=None,
deletedDateTime=None,
roleTemplateId=None,
),
],
)
# Act
result = MicrosoftSSOHandler._get_group_ids_from_graph_api_response(mock_response)
# Assert
assert isinstance(result, list)
assert len(result) == 2
assert "group1" in result
assert "group2" in result
@pytest.mark.asyncio
@pytest.mark.parametrize(
"team_params",
[
# Test case 1: Using DefaultTeamSSOParams
DefaultTeamSSOParams(
max_budget=10, budget_duration="1d", models=["special-gpt-5"]
),
# Test case 2: Using Dict
{"max_budget": 10, "budget_duration": "1d", "models": ["special-gpt-5"]},
],
)
async def test_default_team_params(team_params):
"""
When litellm.default_team_params is set, it should be used to create a new team
"""
# Arrange
litellm.default_team_params = team_params
def mock_jsonify_team_object(db_data):
return db_data
# Mock Prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
mock_prisma.db.litellm_teamtable.create = AsyncMock()
_wire_team_create_tx(mock_prisma)
mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0)
mock_prisma.get_data = AsyncMock(return_value=None)
mock_prisma.jsonify_team_object = MagicMock(side_effect=mock_jsonify_team_object)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
# Act
team_id = str(uuid.uuid4())
await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids(
service_principal_teams=[
MicrosoftServicePrincipalTeam(
principalId=team_id,
principalDisplayName="Test Team",
)
]
)
# Assert
# Verify team was created with correct parameters
mock_prisma.db.litellm_teamtable.create.assert_called_once()
create_call_args = mock_prisma.db.litellm_teamtable.create.call_args.kwargs[
"data"
]
assert create_call_args["team_id"] == team_id
assert create_call_args["team_alias"] == "Test Team"
assert create_call_args["max_budget"] == 10
assert create_call_args["budget_duration"] == "1d"
assert create_call_args["models"] == ["special-gpt-5"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"team_params",
[
DefaultTeamSSOParams(max_budget=10, budget_duration="1d", organization_id="default-org"),
{"max_budget": 10, "budget_duration": "1d", "organization_id": "default-org"},
],
)
async def test_default_team_params_organization_id_reaches_sso_created_team(team_params):
"""The SSO auto-team path builds NewTeamRequest straight from default_team_params,
so a default organization_id must land on the created team row and be validated."""
from litellm.proxy._types import LiteLLM_OrganizationTable
litellm.default_team_params = team_params
mock_prisma = MagicMock()
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
mock_prisma.db.litellm_teamtable.create = AsyncMock()
_wire_team_create_tx(mock_prisma)
mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0)
mock_prisma.get_data = AsyncMock(return_value=None)
mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data)
mock_org = LiteLLM_OrganizationTable(
organization_id="default-org",
budget_id="budget-id",
created_by="admin",
updated_by="admin",
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.management_endpoints.team_endpoints.get_org_object",
AsyncMock(return_value=mock_org),
) as mock_get_org:
team_id = str(uuid.uuid4())
await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids(
service_principal_teams=[
MicrosoftServicePrincipalTeam(
principalId=team_id,
principalDisplayName="Test Team",
)
]
)
mock_prisma.db.litellm_teamtable.create.assert_called_once()
create_call_args = mock_prisma.db.litellm_teamtable.create.call_args.kwargs["data"]
assert create_call_args["organization_id"] == "default-org"
assert mock_get_org.call_args.kwargs["org_id"] == "default-org"
@pytest.mark.asyncio
async def test_create_team_without_default_params():
"""
Test team creation when litellm.default_team_params is None
Should create team with just the basic required fields
"""
# Arrange
litellm.default_team_params = None
def mock_jsonify_team_object(db_data):
return db_data
# Mock Prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
mock_prisma.db.litellm_teamtable.create = AsyncMock()
_wire_team_create_tx(mock_prisma)
mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0)
mock_prisma.get_data = AsyncMock(return_value=None)
mock_prisma.jsonify_team_object = MagicMock(side_effect=mock_jsonify_team_object)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
# Act
team_id = str(uuid.uuid4())
await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids(
service_principal_teams=[
MicrosoftServicePrincipalTeam(
principalId=team_id,
principalDisplayName="Test Team",
)
]
)
# Assert
mock_prisma.db.litellm_teamtable.create.assert_called_once()
create_call_args = mock_prisma.db.litellm_teamtable.create.call_args.kwargs[
"data"
]
assert create_call_args["team_id"] == team_id
assert create_call_args["team_alias"] == "Test Team"
# Should not have any of the optional fields
assert "max_budget" not in create_call_args
assert "budget_duration" not in create_call_args
assert create_call_args["models"] == []
def test_apply_user_info_values_to_sso_user_defined_values():
from litellm.proxy._types import LiteLLM_UserTable, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import (
apply_user_info_values_to_sso_user_defined_values,
)
user_info = LiteLLM_UserTable(
user_id="123",
user_email="test@example.com",
user_role="admin",
)
user_defined_values: SSOUserDefinedValues = {
"models": [],
"user_id": "456",
"user_email": "test@example.com",
"user_role": "admin",
"max_budget": None,
"budget_duration": None,
}
sso_user_defined_values = apply_user_info_values_to_sso_user_defined_values(
user_info=user_info,
user_defined_values=user_defined_values,
)
assert sso_user_defined_values is not None
assert sso_user_defined_values["user_id"] == "123"
def test_apply_user_info_values_to_sso_user_defined_values_with_models():
"""Test that user's models from DB are preserved when they log in via SSO"""
from litellm.proxy._types import LiteLLM_UserTable, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import (
apply_user_info_values_to_sso_user_defined_values,
)
# Simulate existing user with models=['no-default-models'] in DB
user_info = LiteLLM_UserTable(
user_id="123",
user_email="test@example.com",
user_role="admin",
models=["no-default-models"], # User has this set in DB
)
# Simulate SSO login where models defaults to empty list
user_defined_values: SSOUserDefinedValues = {
"models": [], # Empty on SSO login
"user_id": "456",
"user_email": "test@example.com",
"user_role": "admin",
"max_budget": None,
"budget_duration": None,
}
sso_user_defined_values = apply_user_info_values_to_sso_user_defined_values(
user_info=user_info,
user_defined_values=user_defined_values,
)
assert sso_user_defined_values is not None
assert sso_user_defined_values["user_id"] == "123"
# This is the fix: models from DB should be preserved
assert sso_user_defined_values["models"] == ["no-default-models"]
def test_apply_user_info_values_sso_role_takes_precedence():
"""
Test that SSO role takes precedence over DB role.
When Microsoft SSO returns a user_role, it should be used instead of the role stored in the database.
This ensures SSO is the authoritative source for user roles.
"""
from litellm.proxy._types import LiteLLM_UserTable, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import (
apply_user_info_values_to_sso_user_defined_values,
)
user_info = LiteLLM_UserTable(
user_id="123",
user_email="test@example.com",
user_role="internal_user_viewer",
models=["model-1"],
)
user_defined_values: SSOUserDefinedValues = {
"models": [],
"user_id": "456",
"user_email": "test@example.com",
"user_role": "proxy_admin_viewer",
"max_budget": None,
"budget_duration": None,
}
sso_user_defined_values = apply_user_info_values_to_sso_user_defined_values(
user_info=user_info,
user_defined_values=user_defined_values,
)
assert sso_user_defined_values is not None
assert sso_user_defined_values["user_id"] == "123"
assert sso_user_defined_values["user_role"] == "proxy_admin_viewer"
assert sso_user_defined_values["models"] == ["model-1"]
def test_build_sso_user_update_data_with_valid_role():
"""
Test that _build_sso_user_update_data includes role when SSO provides a valid role.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data
sso_result = CustomOpenID(
id="test-user-123",
email="test@example.com",
display_name="Test User",
provider="microsoft",
team_ids=[],
user_role=LitellmUserRoles.PROXY_ADMIN,
)
update_data = _build_sso_user_update_data(
result=sso_result,
user_email="test@example.com",
user_id="test-user-123",
)
assert update_data["user_email"] == "test@example.com"
assert update_data["user_role"] == "proxy_admin"
def test_build_sso_user_update_data_without_role():
"""
Test that _build_sso_user_update_data only includes email when SSO has no role.
"""
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data
sso_result = CustomOpenID(
id="test-user-456",
email="test@example.com",
display_name="Test User",
provider="microsoft",
team_ids=[],
user_role=None,
)
update_data = _build_sso_user_update_data(
result=sso_result,
user_email="test@example.com",
user_id="test-user-456",
)
assert update_data["user_email"] == "test@example.com"
assert "user_role" not in update_data
def test_normalize_email():
"""
Test that normalize_email correctly lowercases email addresses and handles edge cases.
"""
# Test with lowercase email
assert normalize_email("test@example.com") == "test@example.com"
# Test with uppercase email
assert normalize_email("TEST@EXAMPLE.COM") == "test@example.com"
# Test with mixed case email
assert normalize_email("Test.User@Example.COM") == "test.user@example.com"
# Test with None
assert normalize_email(None) is None
# Test with empty string
assert normalize_email("") == ""
def test_build_sso_user_update_data_normalizes_email():
"""
Test that _build_sso_user_update_data normalizes email addresses to lowercase.
"""
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data
sso_result = CustomOpenID(
id="test-user-789",
email="Test.User@Example.COM",
display_name="Test User",
provider="microsoft",
team_ids=[],
user_role=None,
)
update_data = _build_sso_user_update_data(
result=sso_result,
user_email="Test.User@Example.COM",
user_id="test-user-789",
)
# Email should be normalized to lowercase
assert update_data["user_email"] == "test.user@example.com"
assert "user_role" not in update_data
def test_build_sso_user_update_data_fills_empty_user_alias_from_display_name():
"""
An existing SSO user with no alias gets the IdP display name on login.
"""
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data
sso_result = CustomOpenID(
id="S-1-5-21-adfs-user",
email="jane.doe@example.com",
first_name="Jane",
last_name="Doe",
display_name="Doe, Jane",
provider="generic",
team_ids=[],
)
update_data = _build_sso_user_update_data(
result=sso_result,
user_email="jane.doe@example.com",
user_id="S-1-5-21-adfs-user",
existing_user_alias=None,
)
assert update_data == {"user_email": "jane.doe@example.com", "user_alias": "Doe, Jane"}
def test_build_sso_user_update_data_keeps_existing_user_alias():
"""
An alias already stored for the user is never overwritten by the IdP display name.
"""
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data
sso_result = CustomOpenID(
id="S-1-5-21-adfs-user",
email="jane.doe@example.com",
display_name="Doe, Jane",
provider="generic",
team_ids=[],
)
update_data = _build_sso_user_update_data(
result=sso_result,
user_email="jane.doe@example.com",
user_id="S-1-5-21-adfs-user",
existing_user_alias="Admin-set alias",
)
assert update_data == {"user_email": "jane.doe@example.com"}
@pytest.mark.parametrize(
"result, expected_alias",
[
(
CustomOpenID(id="user-1", display_name="Doe, Jane", first_name="Jane", last_name="Doe", team_ids=[]),
"Doe, Jane",
),
(CustomOpenID(id="user-1", first_name="Jane", last_name="Doe", team_ids=[]), "Jane Doe"),
(CustomOpenID(id="user-1", display_name="user-1", first_name="Jane", team_ids=[]), "Jane"),
(CustomOpenID(id="user-1", display_name="user-1", team_ids=[]), None),
(CustomOpenID(id="user-1", display_name=" ", first_name=" Jane ", last_name="Doe", team_ids=[]), "Jane Doe"),
(CustomOpenID(id="user-1", display_name=" ", first_name=" ", team_ids=[]), None),
({"id": "user-1", "display_name": "Dict User", "first_name": None, "last_name": None}, "Dict User"),
(None, None),
],
)
def test_get_sso_user_alias(result: CustomOpenID | dict[str, str | None] | None, expected_alias: str | None):
"""
The alias is the IdP display name unless it is just the user id, then the joined first/last name.
"""
from litellm.proxy.management_endpoints.ui_sso import _get_sso_user_alias
assert _get_sso_user_alias(result) == expected_alias
def test_generic_response_convertor_normalizes_email():
"""
Test that generic_response_convertor normalizes email addresses.
"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
mock_response = {
"preferred_username": "user123",
"email": "Test.User@Example.COM",
"sub": "Test User",
"first_name": "Test",
"last_name": "User",
"provider": "generic",
}
# Mock JWT handler
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = generic_response_convertor(
response=mock_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
role_mappings=None,
)
# Email should be normalized to lowercase
assert result.email == "test.user@example.com"
assert result.id == "user123"
assert result.display_name == "Test User"
@pytest.mark.asyncio
async def test_upsert_sso_user_updates_role_for_existing_user():
"""
Test that upsert_sso_user updates the user role in database when SSO provides a valid role.
When a user's role is updated in the SSO provider (e.g., Azure), the role should be
updated in the LiteLLM database on subsequent logins, not just at initial user creation.
"""
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
# Existing user in DB with old role
existing_user = LiteLLM_UserTable(
user_id="test-user-123",
user_email="test@example.com",
user_role="internal_user",
models=["model-1"],
)
# SSO result with new role (e.g., user was promoted to admin in Azure)
sso_result = CustomOpenID(
id="test-user-123",
email="test@example.com",
display_name="Test User",
provider="microsoft",
team_ids=["team-1"],
user_role=LitellmUserRoles.PROXY_ADMIN,
)
# Act
await SSOAuthenticationHandler.upsert_sso_user(
result=sso_result,
user_info=existing_user,
user_email="test@example.com",
user_defined_values=None,
prisma_client=mock_prisma,
)
# Assert - verify database was updated with both email and role
mock_prisma.db.litellm_usertable.update_many.assert_called_once()
call_args = mock_prisma.db.litellm_usertable.update_many.call_args
assert call_args.kwargs["where"] == {"user_id": "test-user-123"}
assert call_args.kwargs["data"]["user_email"] == "test@example.com"
assert call_args.kwargs["data"]["user_role"] == "proxy_admin"
@pytest.mark.asyncio
async def test_upsert_sso_user_fills_user_alias_for_existing_user():
"""
An existing user row without an alias is updated with the SSO display name on login.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
existing_user = LiteLLM_UserTable(
user_id="S-1-5-21-adfs-user",
user_email="jane.doe@example.com",
user_role="internal_user",
user_alias=None,
)
sso_result = CustomOpenID(
id="S-1-5-21-adfs-user",
email="jane.doe@example.com",
first_name="Jane",
last_name="Doe",
display_name="Doe, Jane",
provider="generic",
team_ids=[],
)
await SSOAuthenticationHandler.upsert_sso_user(
result=sso_result,
user_info=existing_user,
user_email="jane.doe@example.com",
user_defined_values=None,
prisma_client=mock_prisma,
)
mock_prisma.db.litellm_usertable.update_many.assert_called_once_with(
where={"user_id": "S-1-5-21-adfs-user"},
data={"user_email": "jane.doe@example.com", "user_alias": "Doe, Jane"},
)
@pytest.mark.asyncio
async def test_insert_sso_user_sets_user_alias_from_display_name():
"""
A newly created SSO user is inserted with the IdP display name as user_alias.
"""
from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import insert_sso_user
sso_result = CustomOpenID(
id="S-1-5-21-adfs-user",
email="jane.doe@example.com",
first_name="Jane",
last_name="Doe",
display_name="Doe, Jane",
provider="generic",
team_ids=[],
)
user_defined_values: SSOUserDefinedValues = {
"models": [],
"user_id": "S-1-5-21-adfs-user",
"user_email": "jane.doe@example.com",
"max_budget": None,
"user_role": "internal_user",
"budget_duration": None,
}
with patch(
"litellm.proxy.management_endpoints.ui_sso.new_user",
return_value=NewUserResponse(user_id="S-1-5-21-adfs-user", key="sk-xxxxx", teams=None),
) as mock_new_user:
await insert_sso_user(result_openid=sso_result, user_defined_values=user_defined_values)
new_user_request = mock_new_user.call_args.kwargs["data"]
assert new_user_request.user_id == "S-1-5-21-adfs-user"
assert new_user_request.user_email == "jane.doe@example.com"
assert new_user_request.user_alias == "Doe, Jane"
@pytest.mark.asyncio
async def test_upsert_sso_user_does_not_update_invalid_role():
"""
Test that upsert_sso_user does not update the role if SSO provides an invalid role.
If the SSO returns a role that is not a valid LiteLLM role, it should be ignored
and only the email should be updated.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
# Existing user in DB
existing_user = LiteLLM_UserTable(
user_id="test-user-456",
user_email="test@example.com",
user_role="internal_user",
models=[],
)
# SSO result with invalid role - use MagicMock to bypass validation
# This simulates a raw SSO response that has an invalid role string
sso_result = MagicMock()
sso_result.user_role = "invalid_role_not_in_enum"
# Act
await SSOAuthenticationHandler.upsert_sso_user(
result=sso_result,
user_info=existing_user,
user_email="test@example.com",
user_defined_values=None,
prisma_client=mock_prisma,
)
# Assert - verify only email was updated, not role
mock_prisma.db.litellm_usertable.update_many.assert_called_once()
call_args = mock_prisma.db.litellm_usertable.update_many.call_args
assert call_args.kwargs["where"] == {"user_id": "test-user-456"}
assert call_args.kwargs["data"]["user_email"] == "test@example.com"
assert "user_role" not in call_args.kwargs["data"]
@pytest.mark.asyncio
async def test_upsert_sso_user_no_role_in_sso_response():
"""
Test that upsert_sso_user only updates email when SSO response has no role.
When the SSO provider does not return a role, only the email should be updated.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
# Existing user in DB
existing_user = LiteLLM_UserTable(
user_id="test-user-789",
user_email="old@example.com",
user_role="internal_user",
models=[],
)
# SSO result without role
sso_result = CustomOpenID(
id="test-user-789",
email="new@example.com",
display_name="Test User",
provider="microsoft",
team_ids=[],
user_role=None,
)
# Act
await SSOAuthenticationHandler.upsert_sso_user(
result=sso_result,
user_info=existing_user,
user_email="new@example.com",
user_defined_values=None,
prisma_client=mock_prisma,
)
# Assert - verify only email was updated
mock_prisma.db.litellm_usertable.update_many.assert_called_once()
call_args = mock_prisma.db.litellm_usertable.update_many.call_args
assert call_args.kwargs["where"] == {"user_id": "test-user-789"}
assert call_args.kwargs["data"]["user_email"] == "new@example.com"
assert "user_role" not in call_args.kwargs["data"]
def test_get_user_email_and_id_extracts_microsoft_role():
"""
Test that _get_user_email_and_id_from_result extracts user_role from Microsoft SSO.
This ensures Microsoft SSO roles (from app_roles in id_token) are properly
extracted and converted from enum to string.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
result = CustomOpenID(
id="test-user-id",
email="test@example.com",
display_name="Test User",
provider="microsoft",
team_ids=["team-1"],
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
)
parsed = SSOAuthenticationHandler._get_user_email_and_id_from_result(
result=result,
generic_client_id=None,
)
assert parsed.get("user_email") == "test@example.com"
assert parsed.get("user_id") == "test-user-id"
assert parsed.get("user_role") == "proxy_admin_viewer"
@pytest.mark.asyncio
async def test_get_user_info_from_db_user_exists():
"""
Test that get_user_info_from_db finds existing user and calls upsert_sso_user to update.
"""
from litellm.proxy.management_endpoints.ui_sso import get_user_info_from_db
prisma_client = MagicMock()
user_api_key_cache = MagicMock()
proxy_logging_obj = MagicMock()
user_email = "krrishdholakia@gmail.com"
user_defined_values = {
"models": [],
"user_id": "krrishd",
"user_email": "krrishdholakia@gmail.com",
"max_budget": None,
"user_role": None,
"budget_duration": None,
}
args = {
"result": CustomOpenID(
id="krrishd",
email="krrishdholakia@gmail.com",
first_name=None,
last_name=None,
display_name="a3f1c107-04dc-4c93-ae60-7f32eb4b05ce",
picture=None,
provider=None,
team_ids=[],
),
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"proxy_logging_obj": proxy_logging_obj,
"user_email": user_email,
"user_defined_values": user_defined_values,
}
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_object"
) as mock_get_user_object:
await get_user_info_from_db(**args)
mock_get_user_object.assert_called_once()
assert mock_get_user_object.call_args.kwargs["user_id"] == "krrishd"
@pytest.mark.asyncio
async def test_get_user_info_from_db_user_exists_alternate_user_id():
from litellm.proxy.management_endpoints.ui_sso import get_user_info_from_db
prisma_client = MagicMock()
user_api_key_cache = MagicMock()
proxy_logging_obj = MagicMock()
user_email = "krrishdholakia@gmail.com"
user_defined_values = {
"models": [],
"user_id": "krrishd",
"user_email": "krrishdholakia@gmail.com",
"max_budget": None,
"user_role": None,
"budget_duration": None,
}
args = {
"result": CustomOpenID(
id="krrishd",
email="krrishdholakia@gmail.com",
first_name=None,
last_name=None,
display_name="a3f1c107-04dc-4c93-ae60-7f32eb4b05ce",
picture=None,
provider=None,
team_ids=[],
),
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"proxy_logging_obj": proxy_logging_obj,
"user_email": user_email,
"user_defined_values": user_defined_values,
"alternate_user_id": "krrishd-email1234",
}
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_object"
) as mock_get_user_object:
await get_user_info_from_db(**args)
mock_get_user_object.assert_called_once()
assert mock_get_user_object.call_args.kwargs["user_id"] == "krrishd-email1234"
@pytest.mark.asyncio
async def test_get_user_info_from_db_user_not_exists_creates_user():
"""
Test that get_user_info_from_db creates a new user when user doesn't exist in DB.
When get_existing_user_info_from_db returns None, get_user_info_from_db should:
1. Call upsert_sso_user with user_info=None
2. upsert_sso_user should call insert_sso_user to create the user
3. Add user to teams from SSO response
"""
from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import get_user_info_from_db
prisma_client = MagicMock()
user_api_key_cache = MagicMock()
proxy_logging_obj = MagicMock()
user_email = "newuser@example.com"
user_defined_values: SSOUserDefinedValues = {
"models": [],
"user_id": "new-user-123",
"user_email": "newuser@example.com",
"max_budget": None,
"user_role": None,
"budget_duration": None,
}
sso_result = CustomOpenID(
id="new-user-123",
email="newuser@example.com",
first_name="New",
last_name="User",
display_name="New User",
picture=None,
provider="microsoft",
team_ids=["team-1", "team-2"],
)
args = {
"result": sso_result,
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"proxy_logging_obj": proxy_logging_obj,
"user_email": user_email,
"user_defined_values": user_defined_values,
}
# Mock new user response
mock_new_user = NewUserResponse(
user_id="new-user-123",
key="sk-xxxxx",
teams=None,
)
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_existing_user_info_from_db",
return_value=None, # User doesn't exist
) as mock_get_existing,
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.upsert_sso_user",
return_value=mock_new_user,
) as mock_upsert,
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.add_user_to_teams_from_sso_response",
) as mock_add_teams,
):
# Act
user_info = await get_user_info_from_db(**args)
# Assert
# Should try to find user by id
mock_get_existing.assert_called_once()
assert mock_get_existing.call_args.kwargs["user_id"] == "new-user-123"
assert mock_get_existing.call_args.kwargs["user_email"] == "newuser@example.com"
# Should call upsert_sso_user with None user_info
mock_upsert.assert_called_once()
upsert_call_args = mock_upsert.call_args
assert upsert_call_args.kwargs["user_info"] is None
assert upsert_call_args.kwargs["user_email"] == "newuser@example.com"
assert upsert_call_args.kwargs["user_defined_values"] == user_defined_values
# Should add user to teams
mock_add_teams.assert_called_once()
add_teams_call_args = mock_add_teams.call_args
assert add_teams_call_args.kwargs["result"] == sso_result
assert add_teams_call_args.kwargs["user_info"] == mock_new_user
# Should return the new user
assert user_info == mock_new_user
@pytest.mark.asyncio
async def test_get_user_info_from_db_user_exists_updates_user():
"""
Test that get_user_info_from_db updates existing user when user exists in DB.
When get_existing_user_info_from_db returns a user, get_user_info_from_db should:
1. Call upsert_sso_user with the existing user_info
2. upsert_sso_user should update the user in the database
3. Add user to teams from SSO response
"""
from litellm.proxy._types import LiteLLM_UserTable, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import get_user_info_from_db
prisma_client = MagicMock()
user_api_key_cache = MagicMock()
proxy_logging_obj = MagicMock()
user_email = "existing@example.com"
user_defined_values: SSOUserDefinedValues = {
"models": [],
"user_id": "existing-user-456",
"user_email": "existing@example.com",
"max_budget": None,
"user_role": None,
"budget_duration": None,
}
sso_result = CustomOpenID(
id="existing-user-456",
email="existing@example.com",
first_name="Existing",
last_name="User",
display_name="Existing User",
picture=None,
provider="microsoft",
team_ids=["team-3"],
)
# Existing user in DB
existing_user = LiteLLM_UserTable(
user_id="existing-user-456",
user_email="old@example.com",
user_role="internal_user",
models=["gpt-4"],
teams=[],
)
# Updated user after upsert
updated_user = LiteLLM_UserTable(
user_id="existing-user-456",
user_email="existing@example.com", # Updated email
user_role="internal_user",
models=["gpt-4"],
teams=[],
)
args = {
"result": sso_result,
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"proxy_logging_obj": proxy_logging_obj,
"user_email": user_email,
"user_defined_values": user_defined_values,
}
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_existing_user_info_from_db",
return_value=existing_user, # User exists
) as mock_get_existing,
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.upsert_sso_user",
return_value=updated_user,
) as mock_upsert,
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.add_user_to_teams_from_sso_response",
) as mock_add_teams,
):
# Act
user_info = await get_user_info_from_db(**args)
# Assert
# Should find existing user
mock_get_existing.assert_called_once()
assert mock_get_existing.call_args.kwargs["user_id"] == "existing-user-456"
# Should call upsert_sso_user with existing user_info
mock_upsert.assert_called_once()
upsert_call_args = mock_upsert.call_args
assert upsert_call_args.kwargs["user_info"] == existing_user
assert upsert_call_args.kwargs["user_email"] == "existing@example.com"
# Should add user to teams
mock_add_teams.assert_called_once()
add_teams_call_args = mock_add_teams.call_args
assert add_teams_call_args.kwargs["result"] == sso_result
assert add_teams_call_args.kwargs["user_info"] == updated_user
# Should return the updated user
assert user_info == updated_user
@pytest.mark.asyncio
async def test_check_and_update_if_proxy_admin_id():
"""
Test that a user with matching PROXY_ADMIN_ID gets their role updated to admin
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import (
check_and_update_if_proxy_admin_id,
)
# Mock Prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update = AsyncMock()
# Set up test data
test_user_id = "test_admin_123"
test_user_role = "user"
with patch.dict(os.environ, {"PROXY_ADMIN_ID": test_user_id}):
# Act
updated_role = await check_and_update_if_proxy_admin_id(
user_role=test_user_role, user_id=test_user_id, prisma_client=mock_prisma
)
# Assert
assert updated_role == LitellmUserRoles.PROXY_ADMIN.value
mock_prisma.db.litellm_usertable.update.assert_called_once_with(
where={"user_id": test_user_id},
data={"user_role": LitellmUserRoles.PROXY_ADMIN.value},
)
@pytest.mark.asyncio
async def test_check_and_update_if_proxy_admin_id_already_admin():
"""
Test that a user who is already an admin doesn't get their role updated
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import (
check_and_update_if_proxy_admin_id,
)
# Mock Prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update = AsyncMock()
# Set up test data
test_user_id = "test_admin_123"
test_user_role = LitellmUserRoles.PROXY_ADMIN.value
with patch.dict(os.environ, {"PROXY_ADMIN_ID": test_user_id}):
# Act
updated_role = await check_and_update_if_proxy_admin_id(
user_role=test_user_role, user_id=test_user_id, prisma_client=mock_prisma
)
# Assert
assert updated_role == LitellmUserRoles.PROXY_ADMIN.value
mock_prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_get_generic_sso_response_with_additional_headers():
"""
Test that GENERIC_SSO_HEADERS environment variable is correctly processed
and passed to generic_sso.verify_and_process
"""
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
# Arrange
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
generic_client_id = "test_client_id"
redirect_url = "http://test.com/callback"
# Mock response from verify_and_process
mock_sso_response = {
"sub": "test_user_123",
"email": "test@example.com",
"preferred_username": "testuser",
}
# Set up environment variables including GENERIC_SSO_HEADERS
test_env_vars = {
"GENERIC_CLIENT_SECRET": "test_secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
"GENERIC_SSO_HEADERS": "Authorization=Bearer token123, Content-Type=application/json, X-Custom-Header=custom-value",
}
# Expected headers dictionary
expected_headers = {
"Authorization": "Bearer token123",
"Content-Type": "application/json",
"X-Custom-Header": "custom-value",
}
# Mock the SSO provider and its methods
mock_sso_instance = MagicMock()
mock_sso_instance.verify_and_process = AsyncMock(return_value=mock_sso_response)
mock_sso_instance.access_token = (
None # Avoid triggering JWT decode in process_sso_jwt_access_token
)
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(os.environ, test_env_vars):
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
redirect_url=redirect_url,
sso_jwt_handler=None,
)
# Assert
# Verify verify_and_process was called with the correct headers
mock_sso_instance.verify_and_process.assert_called_once_with(
mock_request,
params={"include_client_id": False},
headers=expected_headers,
)
# Verify the result is returned correctly
assert result == mock_sso_response
@pytest.mark.asyncio
async def test_get_generic_sso_response_with_empty_headers():
"""
Test that when GENERIC_SSO_HEADERS is not set, an empty headers dict is passed
"""
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
# Arrange
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
generic_client_id = "test_client_id"
redirect_url = "http://test.com/callback"
mock_sso_response = {
"sub": "test_user_123",
"email": "test@example.com",
"preferred_username": "testuser",
}
# Set up environment variables without GENERIC_SSO_HEADERS
test_env_vars = {
"GENERIC_CLIENT_SECRET": "test_secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
}
# Mock the SSO provider and its methods
mock_sso_instance = MagicMock()
mock_sso_instance.verify_and_process = AsyncMock(return_value=mock_sso_response)
mock_sso_instance.access_token = (
None # Avoid triggering JWT decode in process_sso_jwt_access_token
)
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(os.environ, test_env_vars):
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
redirect_url=redirect_url,
sso_jwt_handler=None,
)
# Assert
# Verify verify_and_process was called with empty headers dict
mock_sso_instance.verify_and_process.assert_called_once_with(
mock_request, params={"include_client_id": False}, headers={}
)
assert result == mock_sso_response
@pytest.mark.asyncio
async def test_get_generic_sso_response_includes_token_claims_when_enabled(monkeypatch):
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_sso_jwt_handler = MagicMock(spec=JWTHandler)
mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"]
mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = []
userinfo = {
"sub": "subject-only",
"groups": ["admins"],
"access_token": "",
}
access_token = pyjwt.encode(
{
"upn": "token-user@example.com",
"email": "token-user@example.com",
"given_name": "Token",
"family_name": "User",
"display_name": "Token User",
},
"test-secret",
algorithm="HS256",
)
mock_sso_instance = MagicMock()
mock_sso_instance.access_token = access_token
mock_sso_instance.id_token = None
def fake_create_provider(*, response_convertor, **_kwargs):
mock_sso_instance.verify_and_process = AsyncMock(
side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object())
)
return MagicMock(return_value=mock_sso_instance)
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret")
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth")
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token")
monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo")
monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true")
monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn")
monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email")
monkeypatch.setenv("GENERIC_USER_FIRST_NAME_ATTRIBUTE", "given_name")
monkeypatch.setenv("GENERIC_USER_LAST_NAME_ATTRIBUTE", "family_name")
monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups")
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider):
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id="test-client",
redirect_url="http://test.com/callback",
sso_jwt_handler=mock_sso_jwt_handler,
)
assert isinstance(result, CustomOpenID)
assert result.id == "token-user@example.com"
assert result.email == "token-user@example.com"
assert result.first_name == "Token"
assert result.last_name == "User"
assert result.display_name == "Token User"
assert result.team_ids == ["team-from-userinfo"]
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
assert received_response is not None
assert "access_token" not in received_response
assert "id_token" not in received_response
assert "refresh_token" not in received_response
@pytest.mark.asyncio
async def test_get_generic_sso_response_does_not_include_token_claims_when_disabled(monkeypatch):
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_sso_jwt_handler = MagicMock(spec=JWTHandler)
mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"]
mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = []
access_token = pyjwt.encode({"upn": "token-user@example.com"}, "test-secret", algorithm="HS256")
userinfo = {"sub": "subject-only", "groups": ["admins"], "access_token": ""}
mock_sso_instance = MagicMock()
mock_sso_instance.access_token = access_token
mock_sso_instance.id_token = None
def fake_create_provider(*, response_convertor, **_kwargs):
mock_sso_instance.verify_and_process = AsyncMock(
side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object())
)
return MagicMock(return_value=mock_sso_instance)
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret")
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth")
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token")
monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo")
monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false")
monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn")
monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email")
monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups")
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider):
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id="test-client",
redirect_url="http://test.com/callback",
sso_jwt_handler=mock_sso_jwt_handler,
)
assert isinstance(result, CustomOpenID)
assert result.id is None
assert result.email is None
assert result.display_name is None
assert result.team_ids == ["team-from-userinfo"]
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
assert received_response == {"sub": "subject-only", "groups": ["admins"]}
def test_merge_sso_token_claims_precedence_and_invalid_tokens():
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import _merge_sso_token_claims
id_token = pyjwt.encode(
{"preferred_username": "id-user", "email": "id@example.com", "id_only": "id-value"},
"test-secret",
algorithm="HS256",
)
access_token = pyjwt.encode(
{"preferred_username": "access-user", "email": "access@example.com", "access_only": "access-value"},
"test-secret",
algorithm="HS256",
)
merged = _merge_sso_token_claims(
userinfo={"preferred_username": "userinfo-user", "email": None, "userinfo_only": "userinfo-value"},
id_token=id_token,
access_token=access_token,
)
assert merged["preferred_username"] == "userinfo-user"
assert merged["email"] == "id@example.com"
assert merged["id_only"] == "id-value"
assert merged["access_only"] == "access-value"
userinfo_only = _merge_sso_token_claims(
userinfo={"sub": "userinfo-user", "email": "userinfo@example.com"},
id_token=pyjwt.encode({}, "test-secret", algorithm="HS256"),
access_token="opaque-access-token",
)
assert userinfo_only == {"sub": "userinfo-user", "email": "userinfo@example.com"}
@pytest.mark.asyncio
async def test_get_generic_sso_response_pkce_merges_token_claims_and_excludes_credentials(monkeypatch):
"""The real PKCE path merges access-token claims and keeps bearer credentials out of received_response.
Only the PKCE verifier cache and the HTTP transport are injected, so
prepare_token_exchange_parameters, _pkce_token_exchange and the claim merge all run for real.
"""
import jwt as pyjwt
from starlette.requests import Request as StarletteRequest
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
access_token = pyjwt.encode(
{"sub": "token-user", "email": "token-user@example.com"}, "test-secret", algorithm="HS256"
)
request = StarletteRequest(
{
"type": "http",
"method": "GET",
"path": "/sso/callback",
"query_string": b"code=test-code&state=test-state",
"headers": [(b"cookie", b"litellm_oauth_state=test-state")],
}
)
pkce_cache = MagicMock(redis_cache=None)
pkce_cache.async_get_cache = AsyncMock(return_value={"code_verifier": "test-code-verifier"})
pkce_cache.async_delete_cache = AsyncMock()
token_endpoint_response = MagicMock(status_code=200)
token_endpoint_response.json.return_value = {
"access_token": access_token,
"id_token": "id-token-secret",
"refresh_token": "refresh-token-secret",
}
token_client = MagicMock()
token_client.post = AsyncMock(return_value=token_endpoint_response)
userinfo_endpoint_response = MagicMock(status_code=200)
userinfo_endpoint_response.json.return_value = {"sub": "userinfo-user"}
userinfo_client = MagicMock()
userinfo_client.get = AsyncMock(return_value=userinfo_endpoint_response)
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret")
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth")
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token")
monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo")
monkeypatch.setenv("GENERIC_CLIENT_USE_PKCE", "true")
monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true")
monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email")
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", pkce_cache),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client",
side_effect=[token_client, userinfo_client],
),
):
result, received_response, _, _ = await get_generic_sso_response(
request=request,
jwt_handler=MagicMock(spec=JWTHandler),
generic_client_id="test-client",
redirect_url="http://test.com/callback",
sso_jwt_handler=None,
)
# The real token exchange ran: it forwarded the cached verifier to the token endpoint.
assert token_client.post.await_args.kwargs["data"]["code_verifier"] == "test-code-verifier"
# UserInfo wins for sub; email exists only on the access token, so the merge must supply it.
assert isinstance(result, CustomOpenID)
assert result.email == "token-user@example.com"
assert received_response == {"sub": "userinfo-user", "email": "token-user@example.com"}
pkce_cache.async_delete_cache.assert_awaited_once_with(key="pkce_verifier:test-state")
@pytest.mark.asyncio
async def test_get_generic_sso_response_ignores_opaque_and_empty_token_claims(monkeypatch):
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
userinfo = {
"preferred_username": "userinfo-user",
"email": "userinfo@example.com",
"sub": "User Info",
}
mock_sso_instance = MagicMock()
mock_sso_instance.access_token = "opaque-access-token"
mock_sso_instance.id_token = pyjwt.encode({}, "test-secret", algorithm="HS256")
def fake_create_provider(*, response_convertor, **_kwargs):
mock_sso_instance.verify_and_process = AsyncMock(
side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object())
)
return MagicMock(return_value=mock_sso_instance)
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret")
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth")
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token")
monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo")
monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true")
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider):
result, received_response, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id="test-client",
redirect_url="http://test.com/callback",
sso_jwt_handler=None,
)
assert isinstance(result, CustomOpenID)
assert result.id == "userinfo-user"
assert result.email == "userinfo@example.com"
assert result.display_name == "User Info"
assert received_response == userinfo
class TestCLISSOCallbackFunction:
"""Test the cli_sso_callback function specifically"""
def test_cli_sso_callback_validation_invalid_key(self):
"""Test CLI SSO callback input validation for invalid key format"""
# Test the validation logic without hitting the database
invalid_keys = [
None,
"",
"invalid-key",
"not-sk-key",
"sk", # too short
]
for invalid_key in invalid_keys:
# This should fail validation before any database operations
# We can test this by checking if the key starts with 'sk-'
if not invalid_key or not invalid_key.startswith("sk-"):
# This would trigger the validation error
assert True # Validation works as expected
class TestCLIPollingFunction:
"""Test the cli_poll_key function specifically"""
def test_cli_poll_key_validation_invalid_format(self):
"""Test CLI polling key format validation"""
# Test key format validation logic
invalid_keys = [
"invalid-key",
"not-sk-key",
"",
"sk", # too short
]
for invalid_key in invalid_keys:
# Validation logic: key must start with 'sk-'
if not invalid_key.startswith("sk-"):
# This would trigger the validation error in the actual function
assert True # Validation works as expected
class TestAuthCallbackRouting:
"""Test the auth_callback function routing logic"""
def test_cli_state_detection_and_routing(self):
"""Test that CLI states are properly detected and would route to CLI callback"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# Test CLI state detection logic
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-test1234567890"
# This mimics the logic in auth_callback
if cli_state and cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the login ID from the state
key_id = cli_state.split(":", 1)[1]
assert key_id == "cli-test1234567890"
else:
pytest.fail("CLI state should have been detected")
def test_non_cli_state_routing(self):
"""Test that non-CLI states don't trigger CLI routing"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
non_cli_states = [
"regular_oauth_state",
"some_random_string",
None,
"",
"not_session_token:something",
]
for state in non_cli_states:
# This mimics the routing logic in auth_callback
should_route_to_cli = state and state.startswith(
f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"
)
assert not should_route_to_cli, f"State '{state}' should not route to CLI"
class TestGoogleLoginCLIIntegration:
"""Test the google_login function with CLI parameters"""
def test_google_login_cli_state_generation(self):
"""Test that google_login generates CLI state when CLI parameters are provided"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Test the CLI state generation logic used in google_login
source = "litellm-cli"
key = "cli-test1234567890"
cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key)
assert cli_state is not None
assert cli_state.startswith("litellm-session-token:")
assert "cli-test1234567890" in cli_state
def test_google_login_no_cli_state_when_missing_params(self):
"""Test that google_login doesn't generate CLI state when CLI parameters are missing"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Test various parameter combinations that shouldn't generate CLI state
test_cases = [
(None, None),
("litellm-cli", None),
(None, "cli-test1234567890"),
("wrong-source", "cli-test1234567890"),
]
for source, key in test_cases:
cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key)
assert (
cli_state is None
), f"CLI state should not be generated for source='{source}', key='{key}'"
class TestSSOHandlerIntegration:
"""Test SSOAuthenticationHandler methods"""
def test_should_use_sso_handler(self):
"""Test the SSO handler detection logic"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Test that SSO handler is used when client IDs are provided
assert (
SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test")
is True
)
assert (
SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test")
is True
)
assert (
SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test")
is True
)
# Test that SSO handler is not used when no client IDs are provided
assert SSOAuthenticationHandler.should_use_sso_handler() is False
assert (
SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False
)
@patch.dict(os.environ, {}, clear=False)
def test_get_redirect_url_for_sso(self):
"""Test the redirect URL generation for SSO"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Remove env vars that override request base_url so the test is
# isolated from local settings.
os.environ.pop("PROXY_BASE_URL", None)
os.environ.pop("SERVER_ROOT_PATH", None)
# Mock request object
mock_request = MagicMock()
mock_request.base_url = "https://test.litellm.ai/"
# Test redirect URL generation
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=mock_request, sso_callback_route="sso/callback"
)
assert redirect_url.startswith("https://test.litellm.ai")
assert "sso/callback" in redirect_url
class TestUISSO_FunctionsExistence:
"""Test that all the new functions exist and are importable"""
def test_cli_sso_callback_exists(self):
"""Test that cli_sso_callback function exists"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
assert callable(cli_sso_callback)
def test_cli_poll_key_exists(self):
"""Test that cli_poll_key function exists"""
from litellm.proxy.management_endpoints.ui_sso import cli_poll_key
assert callable(cli_poll_key)
def test_auth_callback_exists(self):
"""Test that auth_callback function exists"""
from litellm.proxy.management_endpoints.ui_sso import auth_callback
assert callable(auth_callback)
def test_google_login_exists(self):
"""Test that google_login function exists"""
from litellm.proxy.management_endpoints.ui_sso import google_login
assert callable(google_login)
def test_sso_authentication_handler_exists(self):
"""Test that SSOAuthenticationHandler class exists with new methods"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Check that the class exists
assert SSOAuthenticationHandler is not None
# Check that the new _get_cli_state method exists
assert hasattr(SSOAuthenticationHandler, "_get_cli_state")
assert callable(SSOAuthenticationHandler._get_cli_state)
class TestSSOStateHandling:
"""Test the SSO state handling for CLI authentication"""
def test_get_cli_state_valid(self):
"""Test generating CLI state with valid parameters"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli", key="cli-test1234567890"
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "cli-test1234567890" in state
def test_get_cli_state_invalid_source(self):
"""Test generating CLI state with invalid source"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="invalid_source", key="cli-test1234567890"
)
assert state is None
def test_get_cli_state_no_key(self):
"""Test generating CLI state without key"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key=None)
assert state is None
def test_get_cli_state_no_source(self):
"""Test generating CLI state without source"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source=None, key="cli-test1234567890"
)
assert state is None
def test_get_cli_state_ignores_existing_key(self):
"""Test CLI state does not embed an existing key"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli",
key="cli-new-key-1234567890",
existing_key="sk-existing-key-456",
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "cli-new-key-1234567890" in state
assert "sk-existing-key-456" not in state
assert state == "litellm-session-token:cli-new-key-1234567890"
def test_get_cli_state_without_existing_key(self):
"""Test generating CLI state without existing_key"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli", key="cli-new-key-789123456", existing_key=None
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "cli-new-key-789123456" in state
assert state == "litellm-session-token:cli-new-key-789123456"
assert state.count(":") == 1 # Only one colon separator
class TestStateRouting:
"""Test state parameter routing logic"""
def test_cli_state_detection(self):
"""Test detection of CLI state parameters"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# Test CLI state format
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-test1234567890"
assert cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:")
# Test extraction of key from state
key_id = cli_state.split(":", 1)[1]
assert key_id == "cli-test1234567890"
def test_cli_state_parsing_uses_single_login_id(self):
"""Test parsing CLI state with a single login ID"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-key-456123"
# Parse as done in auth_callback
state_parts = cli_state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
assert key_id == "cli-new-key-456123"
def test_cli_state_parsing_without_extra_segments(self):
"""Test parsing CLI state uses a single login ID"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# State format: {PREFIX}:{key}
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-key-999123"
# Parse as done in auth_callback
state_parts = cli_state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
assert key_id == "cli-new-key-999123"
def test_non_cli_state_detection(self):
"""Test detection of non-CLI state parameters"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# Test various non-CLI states
test_states = [
"regular_oauth_state",
"some_random_string",
None,
"",
"not_session_token:something",
]
for state in test_states:
if state:
assert not state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:")
else:
assert state != f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"
class TestHTMLIntegration:
"""Test HTML rendering integration with CLI flow"""
def test_html_render_utils_import(self):
"""Test that HTML render utils can be imported correctly"""
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
)
# Test that function exists and is callable
assert callable(render_cli_sso_success_page)
# Test that it returns expected type
html = render_cli_sso_success_page()
assert isinstance(html, str)
assert len(html) > 0
def test_success_page_instructs_manual_close_without_false_countdown(self):
"""Browsers refuse window.close() on tabs they did not open via window.open()
(the CLI opens the page with webbrowser.open), so a 'closing in 3...' countdown
is a promise the browser usually can't keep and the page gets stuck on
'Closing...'. The page must instead always show the manual-close instruction
and never advertise an auto-close that won't happen.
"""
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
)
html = render_cli_sso_success_page()
assert "You can now close this window and return to your terminal." in html
assert "Closing..." not in html
assert "This window will close in" not in html
class TestCustomUISSO:
"""Test the custom UI SSO sign-in handler functionality"""
def test_enterprise_import_error_handling(self):
"""Test that proper error is raised when enterprise module is not available"""
from unittest.mock import MagicMock, patch
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.example.com/"
# Mock user_custom_ui_sso_sign_in_handler to exist but make enterprise import fail
with patch("litellm.proxy.proxy_server.premium_user", True):
with patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
MagicMock(),
):
with patch.dict(
"sys.modules",
{"litellm_enterprise.proxy.auth.custom_sso_handler": None},
):
# Temporarily mock the google_login function call to test the import error path
async def mock_google_login():
# This mimics the relevant part of google_login that would trigger the import error
try:
from litellm_enterprise.proxy.auth.custom_sso_handler import ( # noqa: F401
EnterpriseCustomSSOHandler,
)
return "success"
except ImportError:
raise ValueError(
"Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise."
)
# Test that the ValueError is raised with the correct message
import pytest
with pytest.raises(
ValueError, match="Enterprise features are not available"
):
asyncio.run(mock_google_login())
@pytest.mark.asyncio
async def test_handle_custom_ui_sso_sign_in_success(self):
"""Test successful custom UI SSO sign-in with valid headers"""
from fastapi_sso.sso.base import OpenID
from litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
)
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
# Mock request with custom headers
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"x-litellm-user-id": "test_user_123",
"x-litellm-user-email": "test@example.com",
"x-forwarded-for": "192.168.1.1",
}
mock_request.base_url = "https://test.litellm.ai/"
mock_request.client.host = "10.0.0.10"
# Mock the custom handler
mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler)
expected_openid = OpenID(
id="test_user_123",
email="test@example.com",
first_name="Test",
last_name="User",
display_name="Test User",
picture=None,
provider="custom",
)
mock_custom_handler.handle_custom_ui_sso_sign_in = AsyncMock(
return_value=expected_openid
)
# Mock the redirect response method
mock_redirect_response = MagicMock()
mock_redirect_response.status_code = 303
with patch("litellm.proxy.proxy_server.premium_user", True):
with patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
mock_custom_handler,
):
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
# Assert
# Verify the custom handler was called with the request
mock_custom_handler.handle_custom_ui_sso_sign_in.assert_called_once_with(
request=mock_request
)
# Verify the redirect response was generated with correct OpenID
mock_get_redirect.assert_called_once_with(
result=expected_openid,
request=mock_request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
# Verify the result is the redirect response
assert result == mock_redirect_response
assert result.status_code == 303
@pytest.mark.asyncio
async def test_handle_custom_ui_sso_sign_in_rejects_untrusted_proxy(self):
"""Custom UI SSO rejects spoofed identity headers from direct clients."""
from litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
)
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"x-litellm-user-id": "admin",
"x-litellm-user-email": "admin@example.com",
}
mock_request.base_url = "https://test.litellm.ai/"
mock_request.client.host = "203.0.113.10"
mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler)
mock_custom_handler.handle_custom_ui_sso_sign_in = AsyncMock()
with patch("litellm.proxy.proxy_server.premium_user", True):
with patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
mock_custom_handler,
):
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with pytest.raises(ValueError, match="not trusted"):
await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
mock_custom_handler.handle_custom_ui_sso_sign_in.assert_not_called()
@pytest.mark.asyncio
async def test_custom_ui_sso_handler_execution_with_real_class(self):
"""
Test that when a user provides a custom class instance, it gets properly executed
and its methods are called with the correct parameters
"""
from fastapi_sso.sso.base import OpenID
from litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
)
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
# Create a real custom handler class instance
class TestCustomSSOHandler(CustomSSOLoginHandler):
def __init__(self):
super().__init__()
self.method_called = False
self.received_request = None
async def handle_custom_ui_sso_sign_in(self, request: Request) -> OpenID:
self.method_called = True
self.received_request = request
# Parse headers like the actual implementation would
request_headers_dict = dict(request.headers)
return OpenID(
id=request_headers_dict.get("x-litellm-user-id", "default_user"),
email=request_headers_dict.get(
"x-litellm-user-email", "default@test.com"
),
first_name="Custom",
last_name="Handler",
display_name="Custom Handler Test",
picture=None,
provider="custom",
)
# Create instance of our test handler
test_handler_instance = TestCustomSSOHandler()
# Mock request with custom headers
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"x-litellm-user-id": "custom_test_user_456",
"x-litellm-user-email": "custom@example.com",
"x-forwarded-for": "10.0.0.1",
}
mock_request.base_url = "https://custom.litellm.ai/"
mock_request.client.host = "10.0.0.20"
# Mock the redirect response method
mock_redirect_response = MagicMock()
mock_redirect_response.status_code = 303
with patch("litellm.proxy.proxy_server.premium_user", True):
with patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
test_handler_instance,
):
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
# Assert that our custom handler was executed
assert test_handler_instance.method_called is True
assert test_handler_instance.received_request == mock_request
# Verify the redirect response was called with the OpenID from our custom handler
mock_get_redirect.assert_called_once()
call_args = mock_get_redirect.call_args.kwargs
# Verify the OpenID object has the expected values from our custom handler
openid_result = call_args["result"]
assert openid_result.id == "custom_test_user_456"
assert openid_result.email == "custom@example.com"
assert openid_result.first_name == "Custom"
assert openid_result.last_name == "Handler"
assert openid_result.display_name == "Custom Handler Test"
assert openid_result.provider == "custom"
# Verify the request and other parameters were passed correctly
assert call_args["request"] == mock_request
assert call_args["received_response"] is None
assert call_args["generic_client_id"] is None
assert call_args["ui_access_mode"] is None
# Verify the result is the redirect response
assert result == mock_redirect_response
assert result.status_code == 303
class TestCLIKeyRegenerationFlow:
"""Test the end-to-end CLI key regeneration flow"""
def test_cli_sso_login_id_validation_restricts_charset(self):
"""Test CLI SSO login IDs only allow the generated character set"""
from litellm.proxy.management_endpoints.ui_sso import (
_is_valid_cli_sso_login_id,
)
assert _is_valid_cli_sso_login_id("cli-test_1234567890")
assert not _is_valid_cli_sso_login_id("cli-session")
assert not _is_valid_cli_sso_login_id("cli-test\n1234567890")
assert not _is_valid_cli_sso_login_id("cli-test\x001234567890")
assert not _is_valid_cli_sso_login_id("sk-test1234567890")
def test_cli_sso_flow_lookup_tells_legacy_clients_to_upgrade(self):
"""Legacy CLIs send self-generated sk-<uuid> login ids; the 400 must say the CLI is outdated"""
from litellm.proxy.management_endpoints.ui_sso import (
_get_cli_sso_flow_or_raise,
)
mock_cache = MagicMock()
mock_cache.get_cache.return_value = None
with pytest.raises(HTTPException) as legacy_exc:
_get_cli_sso_flow_or_raise(
login_id="sk-85c789af-fc21-474c-9dc9-b5d794fe07ec",
cache=mock_cache,
)
assert legacy_exc.value.status_code == 400
assert "out of date" in legacy_exc.value.detail
assert "pip install" in legacy_exc.value.detail
mock_cache.get_cache.assert_not_called()
with pytest.raises(HTTPException) as generic_exc:
_get_cli_sso_flow_or_raise(login_id="not-a-valid-id", cache=mock_cache)
assert generic_exc.value.status_code == 400
assert generic_exc.value.detail == "Invalid CLI login session id"
with pytest.raises(HTTPException) as expired_exc:
_get_cli_sso_flow_or_raise(login_id="cli-test_1234567890", cache=mock_cache)
assert expired_exc.value.status_code == 400
assert "session not found or expired" in expired_exc.value.detail
assert "configure a Redis cache" in expired_exc.value.detail
assert "enable_redis_auth_cache" not in expired_exc.value.detail
def test_cli_sso_flow_is_redis_authoritative_when_redis_attached(self):
"""
When Redis is attached, the CLI SSO flow must be read from and written to
Redis directly, never the in-memory layer. Otherwise the worker that served
/sso/cli/start keeps serving its stale in-memory flow and never sees the
sso_complete/session_data update another worker wrote, which is exactly the
multi-worker failure this fix targets.
"""
from litellm.proxy.management_endpoints.ui_sso import (
CLI_SSO_SESSION_TTL_SECONDS,
_get_cli_sso_flow_cache_key,
_get_cli_sso_flow_or_raise,
_set_cli_sso_flow,
)
login_id = "cli-redis_authoritative_1234567890"
cache_key = _get_cli_sso_flow_cache_key(login_id)
fresh_flow = {"poll_secret_hash": "fresh", "sso_complete": True}
stale_flow = {"poll_secret_hash": "stale", "sso_complete": False}
redis_cache = MagicMock()
redis_cache.get_cache.return_value = fresh_flow
cache = MagicMock()
cache.redis_cache = redis_cache
cache.get_cache.return_value = stale_flow
result = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache)
assert result == fresh_flow
redis_cache.get_cache.assert_called_once_with(key=cache_key)
cache.get_cache.assert_not_called()
_set_cli_sso_flow(login_id=login_id, cache=cache, flow=fresh_flow)
redis_cache.set_cache.assert_called_once_with(
key=cache_key, value=json.dumps(fresh_flow), ttl=CLI_SSO_SESSION_TTL_SECONDS
)
cache.set_cache.assert_not_called()
def test_cli_sso_flow_lookup_treats_an_open_redis_breaker_as_a_miss(self):
"""A Redis read refused by the open circuit breaker is a missing session, not a server error.
The direct Redis read is what keeps the flow authoritative across workers, so the
refusal must not fall back to a possibly stale in-memory copy either.
"""
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
from litellm.proxy.management_endpoints.ui_sso import _get_cli_sso_flow_or_raise
redis_cache = MagicMock()
redis_cache.get_cache.side_effect = RedisCircuitBreakerOpenError("Redis circuit breaker is open")
cache = MagicMock()
cache.redis_cache = redis_cache
cache.get_cache.return_value = {"poll_secret_hash": "stale", "sso_complete": False}
with pytest.raises(HTTPException) as exc_info:
_get_cli_sso_flow_or_raise(login_id="cli-breaker_open_1234567890", cache=cache)
assert exc_info.value.status_code == 400
assert "not found or expired" in exc_info.value.detail
cache.get_cache.assert_not_called()
def test_cli_sso_flow_with_enum_survives_redis_round_trip(self):
"""
RedisCache stores values via str(value) and reads them back through
json.loads/ast.literal_eval. A raw flow dict containing a Python enum
(session_data.user_role after the SSO callback) produces an unparseable
repr, so every worker reading the completed flow from Redis got a
SyntaxError and returned 400 "session not found". The flow must survive
a real Redis serialization round trip.
"""
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import (
_get_cli_sso_flow_or_raise,
_set_cli_sso_flow,
)
login_id = "cli-enum_round_trip_1234567890"
completed_flow = {
"poll_secret_hash": "hash",
"sso_complete": True,
"user_code_verified": False,
"session_data": {
"user_id": "user-1",
"user_role": LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
"models": [],
"teams": ["team-1"],
"team_details": [{"team_id": "team-1", "team_alias": "alias"}],
},
}
redis_store: dict = {}
redis_cache = MagicMock()
redis_cache.set_cache.side_effect = lambda key, value, ttl: redis_store.__setitem__(
key, str(value).encode("utf-8")
)
redis_cache.get_cache.side_effect = lambda key: RedisCache._get_cache_logic(
MagicMock(), redis_store.get(key)
)
cache = MagicMock()
cache.redis_cache = redis_cache
_set_cli_sso_flow(login_id=login_id, cache=cache, flow=completed_flow)
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache)
assert flow["sso_complete"] is True
assert flow["session_data"]["user_role"] == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
assert flow["session_data"]["team_details"] == [{"team_id": "team-1", "team_alias": "alias"}]
@pytest.mark.asyncio
async def test_cli_sso_start_creates_bound_flow(self):
"""Test CLI SSO start creates a polling secret bound flow"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_start,
)
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_cache = MagicMock(redis_cache=None)
mock_cache.increment_cache.return_value = 1
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
result = await cli_sso_start(request=mock_request)
assert result["login_id"].startswith("cli-")
assert result["poll_secret"]
assert result["user_code"]
mock_cache.increment_cache.assert_called_once()
assert mock_cache.increment_cache.call_args.kwargs["ttl"] == 60
mock_cache.set_cache.assert_called_once()
flow_data = mock_cache.set_cache.call_args.kwargs["value"]
assert flow_data["poll_secret_hash"] == _hash_cli_sso_secret(
result["poll_secret"]
)
assert flow_data["user_code_hash"] == _hash_cli_sso_secret(
_normalize_cli_sso_user_code(result["user_code"])
)
assert flow_data["poll_secret_hash"] != result["poll_secret"]
assert flow_data["user_code_hash"] != result["user_code"]
@pytest.mark.asyncio
async def test_cli_sso_start_rate_limits_by_client_ip(self):
"""Test CLI SSO start enforces a coarse per-client rate limit"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_cache = MagicMock(redis_cache=None)
mock_cache.increment_cache.return_value = 31
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_start(request=mock_request)
assert exc_info.value.status_code == 429
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_start_returns_verification_uri_complete_when_enabled(self):
"""Test CLI SSO start returns a verification_uri_complete that round-trips the user_code only when the operator opts in"""
from urllib.parse import parse_qs, urlparse
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_request.base_url = "https://proxy.example.com/"
mock_cache = MagicMock(redis_cache=None)
mock_cache.increment_cache.return_value = 1
with (
patch.dict(
os.environ,
{"PROXY_BASE_URL": "https://proxy.example.com", "SERVER_ROOT_PATH": ""},
),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_cli_sso_verification_uri_complete": True},
),
):
result = await cli_sso_start(request=mock_request)
verification_uri_complete = result["verification_uri_complete"]
parsed = urlparse(verification_uri_complete)
query = parse_qs(parsed.query)
assert parsed.path.endswith("/sso/key/generate")
assert query["source"] == [LITELLM_CLI_SOURCE_IDENTIFIER]
assert query["key"] == [result["login_id"]]
assert query["user_code"] == [result["user_code"]]
@pytest.mark.asyncio
async def test_cli_sso_start_omits_verification_uri_complete_by_default(self):
"""Test CLI SSO start does NOT advertise verification_uri_complete unless the operator enables it (default off)"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_request.base_url = "https://proxy.example.com/"
mock_cache = MagicMock(redis_cache=None)
mock_cache.increment_cache.return_value = 1
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
result = await cli_sso_start(request=mock_request)
assert "verification_uri_complete" not in result
assert result["user_code"]
assert result["login_id"].startswith("cli-")
def test_cli_sso_verification_uri_complete_enabled_reads_general_settings(self):
"""Test the operator opt-in flag is read from general_settings and defaults off"""
from litellm.proxy.management_endpoints.ui_sso import (
_cli_sso_verification_uri_complete_enabled,
)
with patch("litellm.proxy.proxy_server.general_settings", {}):
assert _cli_sso_verification_uri_complete_enabled() is False
with patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_cli_sso_verification_uri_complete": True},
):
assert _cli_sso_verification_uri_complete_enabled() is True
@pytest.mark.asyncio
async def test_google_login_only_threads_user_code_when_enabled(self):
"""Test google_login forwards user_code into the OAuth state only when the operator opt-in is on, dropping it otherwise"""
from litellm.proxy.management_endpoints.ui_sso import google_login
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.example.com/"
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {"poll_secret_hash": "h"}
env_without_sso_providers = {
name: value
for name, value in os.environ.items()
if name not in _SSO_PROVIDER_ENV_VARS
}
async def drive(enabled: bool):
with (
patch.dict(os.environ, env_without_sso_providers, clear=True),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
None,
),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_cli_sso_verification_uri_complete": enabled},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env",
return_value=None,
),
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_redirect_url_for_sso",
return_value="https://proxy.example.com/sso/callback",
),
patch(
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state",
return_value=None,
) as mock_get_cli_state,
):
await google_login(
request=mock_request,
source="litellm-cli",
key="cli-validsessionkey123456",
user_code="WXYZ-2345",
)
assert mock_get_cli_state.called
return mock_get_cli_state.call_args.kwargs["user_code"]
assert await drive(enabled=True) == "WXYZ-2345"
assert await drive(enabled=False) is None
def test_get_cli_state_appends_user_code_for_prefill(self):
"""Test the OAuth state carries the user_code only for the opt-in prefill flow"""
from litellm.constants import (
LITELLM_CLI_SESSION_TOKEN_PREFIX,
LITELLM_CLI_SOURCE_IDENTIFIER,
)
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
manual_state = SSOAuthenticationHandler._get_cli_state(
source=LITELLM_CLI_SOURCE_IDENTIFIER, key="cli-abc123"
)
prefill_state = SSOAuthenticationHandler._get_cli_state(
source=LITELLM_CLI_SOURCE_IDENTIFIER,
key="cli-abc123",
user_code="WXYZ-2345",
)
assert manual_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123"
assert (
prefill_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123:WXYZ-2345"
)
assert (
SSOAuthenticationHandler._get_cli_state(
source="not-cli", key="cli-abc123", user_code="WXYZ-2345"
)
is None
)
def test_get_cli_state_drops_malformed_user_code(self):
"""Test a user_code that is not a server-issued code is dropped before reaching the size-limited OAuth state"""
from litellm.constants import (
LITELLM_CLI_SESSION_TOKEN_PREFIX,
LITELLM_CLI_SOURCE_IDENTIFIER,
)
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
manual_only = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123"
for bad_user_code in ("A" * 4096, "not-a-code", "WXYZ2345", "WXYZ-234", ""):
assert (
SSOAuthenticationHandler._get_cli_state(
source=LITELLM_CLI_SOURCE_IDENTIFIER,
key="cli-abc123",
user_code=bad_user_code,
)
== manual_only
)
def test_is_valid_cli_sso_user_code_matches_generated_format(self):
"""Test the user_code validator accepts a freshly generated code and rejects malformed input"""
from litellm.proxy.management_endpoints.ui_sso import (
_generate_cli_sso_user_code,
_is_valid_cli_sso_user_code,
)
assert _is_valid_cli_sso_user_code(_generate_cli_sso_user_code())
assert _is_valid_cli_sso_user_code("WXYZ-2345")
assert not _is_valid_cli_sso_user_code("WXYZ-2340") # 0 is not in the alphabet
assert not _is_valid_cli_sso_user_code("wxyz-2345")
assert not _is_valid_cli_sso_user_code("WXYZ2345")
assert not _is_valid_cli_sso_user_code("A" * 64)
assert not _is_valid_cli_sso_user_code(None)
def test_cli_state_round_trips_user_code_to_callback_parser(self):
"""Test the callback's state parser recovers login_id and user_code from the state _get_cli_state builds"""
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source=LITELLM_CLI_SOURCE_IDENTIFIER,
key="cli-abc123",
user_code="WXYZ-2345",
)
state_parts = state.split(":", 2)
key_id = state_parts[1] if len(state_parts) > 1 else None
prefill_user_code = state_parts[2] if len(state_parts) > 2 else None
assert key_id == "cli-abc123"
assert prefill_user_code == "WXYZ-2345"
def test_render_cli_sso_verification_page_prefills_user_code(self):
"""Test the verify page pre-fills the user_code input (HTML-escaped) when provided"""
from litellm.proxy.management_endpoints.ui_sso import (
_render_cli_sso_verification_page,
)
html = _render_cli_sso_verification_page(
verify_url="https://proxy.example.com/sso/cli/complete/cli-abc123",
browser_complete_token="browser-token",
prefill_user_code='WXYZ-2345"><script>',
)
assert 'name="user_code"' in html
assert "WXYZ-2345&quot;&gt;&lt;script&gt;" in html
assert '"><script>' not in html
assert "Confirm the verification code below" in html
assert "shown in your terminal" not in html
def test_render_cli_sso_verification_page_omits_value_without_prefill(self):
"""Test the verify page renders the empty manual input when no prefill is provided (backward compatible)"""
from litellm.proxy.management_endpoints.ui_sso import (
_render_cli_sso_verification_page,
)
html = _render_cli_sso_verification_page(
verify_url="https://proxy.example.com/sso/cli/complete/cli-abc123",
browser_complete_token="browser-token",
)
input_line = next(
line for line in html.splitlines() if 'name="user_code"' in line
)
assert "value=" not in input_line
assert "shown in your terminal" in html
@pytest.mark.asyncio
async def test_cli_sso_callback_prefills_user_code_on_verify_page(self):
"""Test the CLI SSO callback threads prefill_user_code into the rendered verify page"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "https://proxy.example.com/"
mock_user_info = LiteLLM_UserTable(
user_id="test-user-123",
user_role="internal_user",
teams=[],
models=[],
)
mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
with (
patch.dict(
os.environ,
{"PROXY_BASE_URL": "https://proxy.example.com", "SERVER_ROOT_PATH": ""},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
return_value=mock_user_info,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
result = await cli_sso_callback(
request=mock_request,
key="cli-session-4567890",
result=mock_sso_result,
prefill_user_code="WXYZ-2345",
)
assert result.status_code == 200
assert 'value="WXYZ-2345"' in result.body.decode()
@pytest.mark.asyncio
async def test_cli_sso_complete_verifies_user_code(self):
"""Test CLI SSO complete marks a session as verified"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"browser_complete_token_hash": _hash_cli_sso_secret("browser-token"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {"user_id": "test-user-123"},
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
return_value="<html>Success</html>",
),
):
result = await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert result.status_code == 200
flow_data = mock_cache.set_cache.call_args.kwargs["value"]
assert flow_data["user_code_verified"] is True
@pytest.mark.asyncio
async def test_cli_sso_complete_requires_callback_token(self):
"""Test CLI SSO complete requires the callback-delivered token"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(return_value=b"user_code=ABCD-EFGH")
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"browser_complete_token_hash": _hash_cli_sso_secret("browser-token"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {"user_id": "test-user-123"},
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert exc_info.value.status_code == 400
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_complete_waits_for_callback_before_token_checks(self):
"""Test CLI SSO complete returns not-ready before verification checks"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert exc_info.value.status_code == 400
assert exc_info.value.detail == "CLI login is not ready"
mock_request.body.assert_not_awaited()
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_callback_stores_session(self):
"""Test CLI SSO callback stores session data in cache for JWT generation"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
# Test data
session_key = "cli-session-4567890"
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id="test-user-123",
user_role="internal_user",
teams=["team1", "team2"],
models=["gpt-4"],
)
# Mock SSO result
mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"}
# Mock cache
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
mock_prisma = MagicMock()
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(
return_value=[
MagicMock(
model_dump=lambda team_id=team_id: {
"team_id": team_id,
"team_alias": team_id,
"models": [],
}
)
for team_id in ("team1", "team2")
]
)
with (
patch.dict(
os.environ,
{
"PROXY_BASE_URL": "https://test.example.com",
"SERVER_ROOT_PATH": "",
},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
return_value=mock_user_info,
),
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
return_value="<html>Success</html>",
),
):
# Act
result = await cli_sso_callback(
request=mock_request,
key=session_key,
result=mock_sso_result,
)
# Assert - verify session was stored in cache
mock_cache.set_cache.assert_called_once()
call_args = mock_cache.set_cache.call_args
# Verify cache key format
assert "cli_sso_session:" in call_args.kwargs["key"]
assert session_key in call_args.kwargs["key"]
# Verify session data structure
flow_data = call_args.kwargs["value"]
session_data = flow_data["session_data"]
assert flow_data["sso_complete"] is True
assert flow_data["user_code_verified"] is False
assert isinstance(flow_data["browser_complete_token_hash"], str)
assert session_data["user_id"] == "test-user-123"
assert session_data["user_role"] == "internal_user"
assert session_data["teams"] == ["team1", "team2"]
assert session_data["models"] == ["gpt-4"]
# Verify TTL
assert call_args.kwargs["ttl"] == 600
assert result.status_code == 200
# Verify response contains success message (response is HTML)
assert result.body is not None
assert (
'action="https://test.example.com/sso/cli/complete/cli-session-4567890"'
in result.body.decode()
)
@pytest.mark.asyncio
async def test_cli_poll_key_returns_teams_for_selection(self):
"""Test CLI poll endpoint returns teams for user selection when multiple teams exist"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
# Test data
session_key = "cli-session-789123"
session_data = {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": ["team-a", "team-b", "team-c"],
"models": ["gpt-4"],
}
# Mock cache
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
# Act - First poll without team_id
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
# Assert - should return teams list for selection
assert result["status"] == "ready"
assert result["requires_team_selection"] is True
assert result["user_id"] == "test-user-456"
assert result["teams"] == ["team-a", "team-b", "team-c"]
assert "key" not in result # JWT should not be generated yet
# Verify session was NOT deleted
mock_cache.delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_poll_key_requires_poll_secret(self):
"""Test CLI poll endpoint rejects callers without the polling secret"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
},
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
with pytest.raises(HTTPException) as exc_info:
await cli_poll_key(key_id="cli-session-789123", team_id=None)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_cli_poll_key_waits_for_user_code_verification(self):
"""Test CLI poll endpoint stays pending until user code verification"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
},
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
result = await cli_poll_key(
key_id="cli-session-789123",
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result == {"status": "pending"}
@pytest.mark.asyncio
async def test_auth_callback_routes_to_cli(self):
"""Test that auth_callback properly routes CLI requests"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
from litellm.proxy.management_endpoints.ui_sso import auth_callback
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"code": "some-auth-code"}
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-session-key-456"
# Mock the CLI callback and required proxy server components
mock_result = {"user_id": "test-user", "email": "test@example.com"}
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.cli_sso_callback"
) as mock_cli_callback,
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.master_key", "test-master-key"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch.dict(os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True),
patch(
"litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response",
return_value=mock_result,
),
):
mock_cli_callback.return_value = MagicMock()
# Act
await auth_callback(request=mock_request, state=cli_state)
mock_cli_callback.assert_called_once_with(
request=mock_request,
key="cli-new-session-key-456",
prefill_user_code=None,
result=mock_result,
received_response=None,
sso_assertion=None,
)
@pytest.mark.asyncio
async def test_auth_callback_forwards_prefill_user_code_from_state(self):
"""Test auth_callback recovers the user_code from the state and forwards it for prefill"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
from litellm.proxy.management_endpoints.ui_sso import auth_callback
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"code": "some-auth-code"}
cli_state = (
f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-session-key-456:WXYZ-2345"
)
mock_result = {"user_id": "test-user", "email": "test@example.com"}
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.cli_sso_callback"
) as mock_cli_callback,
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.master_key", "test-master-key"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch.dict(os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True),
patch(
"litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response",
return_value=mock_result,
),
):
mock_cli_callback.return_value = MagicMock()
await auth_callback(request=mock_request, state=cli_state)
mock_cli_callback.assert_called_once_with(
request=mock_request,
key="cli-new-session-key-456",
prefill_user_code="WXYZ-2345",
result=mock_result,
received_response=None,
sso_assertion=None,
)
def test_get_redirect_url_does_not_include_existing_key_in_url(self):
"""Test that redirect URL generation does NOT include existing_key in URL"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.litellm.ai/"
with patch(
"litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"
):
# Test with existing_key - should NOT be in URL
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=mock_request,
sso_callback_route="sso/callback",
existing_key="sk-existing-123",
)
# existing_key should NOT be in the URL
assert "https://test.litellm.ai/sso/callback" == redirect_url
assert "existing_key" not in redirect_url
def test_get_redirect_url_without_existing_key(self):
"""Test that redirect URL generation works without existing_key parameter"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.litellm.ai/"
with patch(
"litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"
):
# Test without existing_key
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=mock_request, sso_callback_route="sso/callback"
)
assert "https://test.litellm.ai/sso/callback" == redirect_url
@pytest.mark.asyncio
async def test_cli_poll_key_generates_jwt_with_team(self):
"""Test CLI poll endpoint generates JWT when team_id is provided"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
# Test data
session_key = "cli-session-999123"
selected_team = "team-b"
session_data = {
"user_id": "test-user-789",
"user_role": "internal_user",
"teams": ["team-a", "team-b", "team-c"],
"team_details": [
{"team_id": "team-a", "team_alias": "Team A", "team_models": []},
{"team_id": "team-b", "team_alias": "Team B", "team_models": []},
{"team_id": "team-c", "team_alias": "Team C", "team_models": []},
],
"models": ["gpt-4"],
"user_email": "test@example.com",
}
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id="test-user-789",
user_role="internal_user",
teams=["team-a", "team-b", "team-c"],
models=["gpt-4"],
)
# Mock cache
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.token"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.prisma_client"),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
) as mock_get_jwt,
patch(
"litellm.proxy.auth.auth_checks.get_user_object",
new=AsyncMock(return_value=mock_user_info),
),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(side_effect=Exception("no team")),
),
):
# Act - Second poll with team_id
result = await cli_poll_key(
key_id=session_key,
team_id=selected_team,
x_litellm_cli_poll_secret="poll-secret",
)
# Assert - should return JWT
assert result["status"] == "ready"
assert result["key"] == mock_jwt_token
assert result["user_id"] == "test-user-789"
assert result["team_id"] == selected_team
assert result["teams"] == ["team-a", "team-b", "team-c"]
# Verify JWT was generated with correct team and no budget cap
# (team lookup failed, but team_id is set, so fallback cap must not apply)
mock_get_jwt.assert_called_once()
jwt_call_args = mock_get_jwt.call_args
assert jwt_call_args.kwargs["team_id"] == selected_team
assert jwt_call_args.kwargs["team_alias"] == "Team B"
assert jwt_call_args.kwargs["max_budget"] is None
# Verify session was deleted after JWT generation
mock_cache.delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_fetch_cli_sso_team_details_projects_team_grants(self):
"""The cached team detail must carry the team's model grants.
The projection used to drop everything except team_id/team_alias, so the
minted CLI token had no team_models and no team_model_aliases to snapshot.
The joined alias table is stored JSON-encoded, so it has to be decoded here
too, otherwise alias lookup at request time is a substring match on a string.
"""
from litellm.proxy.management_endpoints.ui_sso import (
fetch_cli_sso_team_details,
)
team_row = MagicMock()
team_row.model_dump.return_value = {
"team_id": "team-a",
"team_alias": "Team A",
"models": ["claude-sonnet-4-5", "gpt-4.1"],
"litellm_model_table": {
"id": 7,
"model_aliases": json.dumps({"team-fast": "gpt-4.1-mini"}),
"created_by": "admin",
"updated_by": "admin",
},
}
find_many = AsyncMock(return_value=[team_row])
prisma_client = MagicMock()
prisma_client.db.litellm_teamtable.find_many = find_many
details = await fetch_cli_sso_team_details(
prisma_client=prisma_client, teams=["team-a"]
)
assert find_many.await_args.kwargs["include"] == {"litellm_model_table": True}
assert [detail.model_dump() for detail in details] == [
{
"team_id": "team-a",
"team_alias": "Team A",
"team_models": ("claude-sonnet-4-5", "gpt-4.1"),
"team_model_aliases": {"team-fast": "gpt-4.1-mini"},
}
]
@pytest.mark.asyncio
async def test_fetch_cli_sso_team_details_separates_lookup_failure_from_no_teams(self):
"""A failed lookup must not look like a team that resolved to nothing.
Both used to return [], so a database blip was indistinguishable from a real
answer. The callback needs them apart: a blip has to fail the login, while a
real empty answer means the team rows are genuinely gone.
"""
from litellm.proxy.management_endpoints.ui_sso import (
fetch_cli_sso_team_details,
)
failing_client = MagicMock()
failing_client.db.litellm_teamtable.find_many = AsyncMock(
side_effect=Exception("connection reset")
)
assert (
await fetch_cli_sso_team_details(
prisma_client=failing_client, teams=["team-a"]
)
is None
)
empty_client = MagicMock()
empty_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
assert (
await fetch_cli_sso_team_details(
prisma_client=empty_client, teams=["team-a"]
)
== ()
)
@pytest.mark.asyncio
async def test_cli_poll_key_mints_jwt_with_selected_team_grants(self):
"""The selected team's grants must reach the mint, not just its alias."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_data = {
"user_id": "grants-user",
"user_role": "internal_user",
"teams": ["team-a", "team-b"],
"team_details": [
{
"team_id": "team-a",
"team_alias": "Team A",
"team_models": ["gpt-4.1"],
"team_model_aliases": {"a-fast": "gpt-4.1-mini"},
},
{
"team_id": "team-b",
"team_alias": "Team B",
"team_models": ["claude-sonnet-4-5"],
"team_model_aliases": {"b-fast": "claude-haiku-4-5"},
},
],
"models": ["personal-only"],
"user_email": "grants@example.com",
}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
with (
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value="minted-token",
) as mock_get_jwt,
):
result = await cli_poll_key(
key_id="cli-session-grants",
team_id="team-b",
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
kwargs = mock_get_jwt.call_args.kwargs
assert kwargs["team_id"] == "team-b"
assert kwargs["team_alias"] == "Team B"
assert kwargs["team_models"] == ("claude-sonnet-4-5",)
assert kwargs["team_model_aliases"] == {"b-fast": "claude-haiku-4-5"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"team_details",
[
pytest.param(None, id="detail_fetch_failed"),
pytest.param(
[{"team_id": "team-other", "team_models": []}], id="selected_team_absent"
),
pytest.param(
[{"team_id": "team-a", "team_alias": "Team A"}],
id="legacy_detail_without_grants",
),
],
)
async def test_cli_poll_key_refuses_to_mint_when_team_grants_are_unknown(
self, team_details
):
"""An unknown team grant must never be minted as an empty one.
get_complete_model_list falls through to the whole proxy model list when both
the key allowlist and the team allowlist are empty, and team-bound tokens carry
an empty key allowlist by design. So minting an unresolved team as empty would
hand a team-bound CLI session every model on the proxy.
"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": {
"user_id": "grants-user",
"user_role": "internal_user",
"teams": ["team-a"],
"team_details": team_details,
"models": ["personal-only"],
"user_email": "grants@example.com",
},
}
with (
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value="minted-token",
) as mock_get_jwt,
):
with pytest.raises(HTTPException) as exc_info:
await cli_poll_key(
key_id="cli-session-grants",
team_id="team-a",
x_litellm_cli_poll_secret="poll-secret",
)
assert exc_info.value.status_code == 500
assert "team-a" in str(exc_info.value.detail)
mock_get_jwt.assert_not_called()
mock_cache.delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_poll_key_mints_teamless_session_without_team_grants(self):
"""A user with no team still mints, keeping their personal allowlist in the key slot."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": {
"user_id": "teamless-user",
"user_role": "internal_user",
"teams": [],
"team_details": [],
"models": ["personal-only"],
"user_email": "teamless@example.com",
},
}
with (
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value="minted-token",
) as mock_get_jwt,
):
result = await cli_poll_key(
key_id="cli-session-teamless",
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
kwargs = mock_get_jwt.call_args.kwargs
assert kwargs["team_id"] is None
assert kwargs["team_models"] == ()
assert kwargs["user_info"].models == ["personal-only"]
@pytest.mark.asyncio
async def test_cli_poll_key_does_not_cap_session_when_user_has_budget(self):
"""A user with a configured budget must not get the max_ui_session_budget fallback cap."""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_data = {
"user_id": "budgeted-user",
"user_role": "internal_user",
"teams": [],
"team_details": [],
"models": ["gpt-4"],
"user_email": "budgeted@example.com",
}
mock_user_info = LiteLLM_UserTable(
user_id="budgeted-user",
user_role="internal_user",
teams=[],
models=["gpt-4"],
max_budget=100.0,
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiJ9.budgeted.token"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.prisma_client"),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
) as mock_get_jwt,
patch(
"litellm.proxy.auth.auth_checks.get_user_object",
new=AsyncMock(return_value=mock_user_info),
),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(
side_effect=AssertionError("team lookup must be skipped")
),
),
):
result = await cli_poll_key(
key_id="cli-session-budgeted",
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
mock_get_jwt.assert_called_once()
assert mock_get_jwt.call_args.kwargs["max_budget"] is None
@pytest.mark.asyncio
async def test_cli_poll_key_does_not_cap_session_even_without_user_or_team_budget(self):
"""Regression: a CLI session token must not inherit the UI chat-pane budget
(max_ui_session_budget). Even when the user and team have no budget of their
own, the minted token carries max_budget=None and is governed only by the
real user/team budgets at request time."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_data = {
"user_id": "unbudgeted-user",
"user_role": "internal_user",
"teams": ["team-x"],
"team_details": [{"team_id": "team-x", "team_alias": "Team X", "team_models": []}],
"models": ["gpt-4"],
"user_email": "unbudgeted@example.com",
}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiJ9.unbudgeted.token"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
) as mock_get_jwt,
):
result = await cli_poll_key(
key_id="cli-session-unbudgeted",
team_id="team-x",
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
mock_get_jwt.assert_called_once()
assert mock_get_jwt.call_args.kwargs["max_budget"] is None
class TestGetAppRolesFromIdToken:
"""Test the get_app_roles_from_id_token method"""
def test_roles_picked_when_app_roles_not_exists(self):
"""Test that 'roles' is picked when 'app_roles' doesn't exist"""
# Create a token with only 'roles' claim
token_payload = {
"sub": "user123",
"email": "test@example.com",
"roles": ["Admin", "User", "Developer"],
}
# Create a mock JWT token
mock_token = "mock.jwt.token"
with patch("jwt.decode", return_value=token_payload) as mock_jwt_decode:
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert
assert result == ["Admin", "User", "Developer"]
mock_jwt_decode.assert_called_once_with(
mock_token, options={"verify_signature": False}
)
def test_app_roles_picked_when_both_exist(self):
"""Test that 'app_roles' takes precedence when both 'app_roles' and 'roles' exist"""
# Create a token with both 'app_roles' and 'roles' claims
token_payload = {
"sub": "user123",
"email": "test@example.com",
"app_roles": ["AppAdmin", "AppUser"],
"roles": ["RoleAdmin", "RoleUser"],
}
mock_token = "mock.jwt.token"
with patch("jwt.decode", return_value=token_payload):
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert - app_roles should be picked, not roles
assert result == ["AppAdmin", "AppUser"]
def test_roles_picked_when_app_roles_is_empty(self):
"""Test that 'roles' is picked when 'app_roles' exists but is empty"""
# Create a token with empty 'app_roles' and populated 'roles'
token_payload = {
"sub": "user123",
"email": "test@example.com",
"app_roles": [],
"roles": ["Admin", "User"],
}
mock_token = "mock.jwt.token"
with patch("jwt.decode", return_value=token_payload):
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert - roles should be picked since app_roles is empty
assert result == ["Admin", "User"]
def test_empty_list_when_neither_exists(self):
"""Test that empty list is returned when neither 'app_roles' nor 'roles' exist"""
# Create a token without roles claims
token_payload = {"sub": "user123", "email": "test@example.com"}
mock_token = "mock.jwt.token"
with patch("jwt.decode", return_value=token_payload):
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert
assert result == []
def test_empty_list_when_no_token_provided(self):
"""Test that empty list is returned when no token is provided"""
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(None)
# Assert
assert result == []
def test_empty_list_when_roles_not_a_list(self):
"""Test that empty list is returned when roles is not a list"""
# Create a token with non-list roles
token_payload = {
"sub": "user123",
"email": "test@example.com",
"roles": "Admin", # String instead of list
}
mock_token = "mock.jwt.token"
with patch("jwt.decode", return_value=token_payload):
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert
assert result == []
def test_error_handling_on_jwt_decode_exception(self):
"""Test that exceptions during JWT decode are handled gracefully"""
mock_token = "invalid.jwt.token"
with patch("jwt.decode", side_effect=Exception("Invalid token")):
# Act
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token)
# Assert - should return empty list on error
assert result == []
class TestProcessSSOJWTAccessToken:
"""Test the process_sso_jwt_access_token helper function"""
@pytest.fixture
def mock_jwt_handler(self):
"""Create a mock JWT handler for testing"""
mock_handler = MagicMock(spec=JWTHandler)
mock_handler.get_team_ids_from_jwt.return_value = ["team1", "team2", "team3"]
return mock_handler
@pytest.fixture
def sample_jwt_token(self):
"""Create a sample JWT token string"""
return "test-jwt-token-header.payload.signature"
@pytest.fixture
def sample_jwt_payload(self):
"""Create a sample JWT payload"""
return {
"sub": "1234567890",
"name": "John Doe",
"iat": 1516239022,
"groups": ["team1", "team2", "team3"],
}
def test_process_sso_jwt_access_token_with_existing_team_ids(
self, mock_jwt_handler, sample_jwt_token
):
"""Test that existing team IDs are not overwritten"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
# Create a result object with existing team_ids
existing_team_ids = ["existing_team1", "existing_team2"]
result = CustomOpenID(
id="test_user",
email="test@example.com",
first_name="Test",
last_name="User",
display_name="Test User",
provider="generic",
team_ids=existing_team_ids,
)
with patch("jwt.decode") as mock_jwt_decode:
# Act
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Assert
# JWT should still be decoded
mock_jwt_decode.assert_called_once()
# But team IDs should NOT be extracted since they already exist
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
# Existing team IDs should remain unchanged
assert result.team_ids == existing_team_ids
def test_process_sso_jwt_access_token_with_dict_result(
self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload
):
"""Test processing with a dictionary result object"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
# Create a dictionary result without team_ids
result = {"id": "test_user", "email": "test@example.com", "name": "Test User"}
with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode:
# Act
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Assert
mock_jwt_decode.assert_called_once_with(
sample_jwt_token, options={"verify_signature": False}
)
mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(
sample_jwt_payload
)
# Verify team_ids was added to the dict as a key
assert "team_ids" in result
assert result["team_ids"] == ["team1", "team2", "team3"]
def test_process_sso_jwt_access_token_with_dict_existing_team_ids(
self, mock_jwt_handler, sample_jwt_token
):
"""Test that existing team IDs in dictionary are not overwritten"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
# Create a dictionary result with existing team_ids
existing_team_ids = ["dict_team1", "dict_team2"]
result = {
"id": "test_user",
"email": "test@example.com",
"name": "Test User",
"team_ids": existing_team_ids,
}
with patch("jwt.decode") as mock_jwt_decode:
# Act
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Assert
# JWT should still be decoded
mock_jwt_decode.assert_called_once()
# But team IDs should NOT be extracted since they already exist
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
# Existing team IDs should remain unchanged
assert result["team_ids"] == existing_team_ids
def test_process_sso_jwt_access_token_no_access_token(self, mock_jwt_handler):
"""Test that nothing happens when access token is None or empty"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[])
# Test with None access token
with patch("jwt.decode") as mock_jwt_decode:
process_sso_jwt_access_token(
access_token_str=None, sso_jwt_handler=mock_jwt_handler, result=result
)
# Assert nothing was processed
mock_jwt_decode.assert_not_called()
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
assert result.team_ids == []
# Test with empty string access token
with patch("jwt.decode") as mock_jwt_decode:
process_sso_jwt_access_token(
access_token_str="", sso_jwt_handler=mock_jwt_handler, result=result
)
# Assert nothing was processed
mock_jwt_decode.assert_not_called()
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
assert result.team_ids == []
def test_process_sso_jwt_access_token_no_result(
self, mock_jwt_handler, sample_jwt_token
):
"""Test that nothing happens when result is None"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
with patch("jwt.decode") as mock_jwt_decode:
# Act
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=None,
)
# Assert nothing was processed
mock_jwt_decode.assert_not_called()
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
def test_process_sso_jwt_access_token_non_decode_exception_propagates(
self, mock_jwt_handler, sample_jwt_token
):
"""Test that non-DecodeError JWT exceptions still propagate up."""
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[])
with patch(
"jwt.decode", side_effect=pyjwt.exceptions.InvalidKeyError("Invalid key")
) as mock_jwt_decode:
with pytest.raises(pyjwt.exceptions.InvalidKeyError, match="Invalid key"):
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
mock_jwt_decode.assert_called_once()
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
def test_process_sso_jwt_access_token_empty_team_ids_from_jwt(
self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload
):
"""Test processing when JWT handler returns empty team IDs"""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
# Configure mock to return empty team IDs
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[])
with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode:
# Act
process_sso_jwt_access_token(
access_token_str=sample_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Assert
mock_jwt_decode.assert_called_once()
mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(
sample_jwt_payload
)
# Even empty team IDs should be set
assert result.team_ids == []
def test_process_sso_jwt_access_token_with_opaque_token(self, mock_jwt_handler):
"""Test that opaque (non-JWT) access tokens are handled gracefully without raising."""
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
result = CustomOpenID(
id="test_user",
email="test@example.com",
first_name="Test",
last_name="User",
display_name="Test User",
provider="generic",
team_ids=["existing_team"],
user_role=None,
)
# Opaque tokens like those from Logto are short random strings, not JWTs
opaque_token = "uTxyjXbS_random_opaque_token_string"
# Should NOT raise - opaque tokens should be silently skipped
process_sso_jwt_access_token(
access_token_str=opaque_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Result should be untouched
mock_jwt_handler.get_team_ids_from_jwt.assert_not_called()
assert result.team_ids == ["existing_team"]
assert result.user_role is None
def test_process_sso_jwt_access_token_real_jwt_with_role_and_teams(
self, mock_jwt_handler
):
"""Test that a real JWT containing role and team fields is correctly processed."""
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
payload = {
"sub": "user123",
"email": "admin@example.com",
"role": "proxy_admin",
"groups": ["team_alpha", "team_beta"],
}
real_jwt_token = pyjwt.encode(payload, "test-secret", algorithm="HS256")
mock_jwt_handler.get_team_ids_from_jwt.return_value = [
"team_alpha",
"team_beta",
]
result = CustomOpenID(
id="user123",
email="admin@example.com",
first_name="Admin",
last_name="User",
display_name="Admin User",
provider="generic",
team_ids=[],
user_role=None,
)
process_sso_jwt_access_token(
access_token_str=real_jwt_token,
sso_jwt_handler=mock_jwt_handler,
result=result,
)
# Team IDs should be extracted via sso_jwt_handler
mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(payload)
assert result.team_ids == ["team_alpha", "team_beta"]
# Role should be extracted from the "role" field in the JWT
from litellm.proxy._types import LitellmUserRoles
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_process_sso_jwt_access_token_real_jwt_without_role_and_teams(self):
"""Test that a real JWT without role/team fields leaves result unchanged."""
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import (
process_sso_jwt_access_token,
)
payload = {
"sub": "user456",
"email": "plain@example.com",
"iat": 1700000000,
}
real_jwt_token = pyjwt.encode(payload, "test-secret", algorithm="HS256")
result = CustomOpenID(
id="user456",
email="plain@example.com",
first_name="Plain",
last_name="User",
display_name="Plain User",
provider="generic",
team_ids=[],
user_role=None,
)
# No sso_jwt_handler, no role/team fields in JWT
process_sso_jwt_access_token(
access_token_str=real_jwt_token,
sso_jwt_handler=None,
result=result,
)
# Nothing should be modified
assert result.team_ids == []
assert result.user_role is None
@pytest.mark.asyncio
async def test_get_ui_settings_includes_api_doc_base_url():
"""Ensure the UI settings endpoint surfaces the optional API doc override."""
from fastapi import Request
from litellm.proxy.management_endpoints.ui_sso import get_ui_settings
mock_request = Request(
scope={
"type": "http",
"headers": [],
"method": "GET",
"scheme": "http",
"server": ("testserver", 80),
"path": "/sso/get/ui_settings",
"query_string": b"",
}
)
with patch.dict(
os.environ,
{
"LITELLM_UI_API_DOC_BASE_URL": "https://custom.docs",
},
):
response = await get_ui_settings(mock_request)
assert response["LITELLM_UI_API_DOC_BASE_URL"] == "https://custom.docs"
class TestGenericResponseConvertorNestedAttributes:
"""Test generic_response_convertor with nested attribute paths"""
def test_generic_response_convertor_with_nested_attributes(self):
"""
Test that generic_response_convertor handles nested attributes with dotted notation
like "attributes.userId"
"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
# Mock JWT handler
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
# Payload with nested attributes structure
nested_payload = {
"sub": "user-sub-123",
"service": "test-service",
"auth_time": 1234567890,
"attributes": {
"given_name": "John",
"oauthClientId": "client-123",
"family_name": "Doe",
"userId": "nested-user-456",
"email": "john.doe@example.com",
},
"id": "top-level-id-789",
"client_id": "client-abc",
}
# Test with nested user ID attribute
with patch.dict(
os.environ,
{
"GENERIC_USER_ID_ATTRIBUTE": "attributes.userId",
"GENERIC_USER_EMAIL_ATTRIBUTE": "attributes.email",
"GENERIC_USER_FIRST_NAME_ATTRIBUTE": "attributes.given_name",
"GENERIC_USER_LAST_NAME_ATTRIBUTE": "attributes.family_name",
"GENERIC_USER_DISPLAY_NAME_ATTRIBUTE": "sub",
},
):
# Act
result = generic_response_convertor(
response=nested_payload,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
)
# Assert
assert isinstance(result, CustomOpenID)
# Note: The current implementation uses response.get() which doesn't support
# dotted notation for nested attributes. This test documents the current behavior.
# If nested attribute support is needed, the implementation would need to be updated
# to handle dotted paths like "attributes.userId"
# Current behavior: returns None for nested paths
# Expected behavior with current implementation (no nested path support):
assert result.id == "nested-user-456"
assert (
result.email == "john.doe@example.com"
) # Can't access "attributes.email" with .get()
assert (
result.first_name == "John"
) # Can't access "attributes.given_name" with .get()
assert (
result.last_name == "Doe"
) # Can't access "attributes.family_name" with .get()
assert result.display_name == "user-sub-123" # Top-level attribute works
class TestGenericResponseConvertorUserRole:
"""Test generic_response_convertor user role extraction from SSO token"""
def test_generic_response_convertor_extracts_valid_user_role(self):
"""
Test that generic_response_convertor extracts a valid LiteLLM user role
from the SSO token using the GENERIC_USER_ROLE_ATTRIBUTE env var.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
sso_response = {
"preferred_username": "testuser",
"email": "test@example.com",
"sub": "Test User",
"role": "proxy_admin",
}
with patch.dict(
os.environ,
{"GENERIC_USER_ROLE_ATTRIBUTE": "role"},
):
result = generic_response_convertor(
response=sso_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
)
assert isinstance(result, CustomOpenID)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_generic_response_convertor_ignores_invalid_user_role(self):
"""
Test that generic_response_convertor ignores invalid role values
and sets user_role to None.
"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
sso_response = {
"preferred_username": "testuser",
"email": "test@example.com",
"role": "invalid_role_value",
}
with patch.dict(
os.environ,
{"GENERIC_USER_ROLE_ATTRIBUTE": "role"},
):
result = generic_response_convertor(
response=sso_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
)
assert isinstance(result, CustomOpenID)
assert result.user_role is None
class TestGetGenericSSORedirectParams:
"""Test _get_generic_sso_redirect_params state parameter priority handling"""
def test_state_priority_cli_state_provided(self):
"""
Test that CLI state takes highest priority when provided
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
cli_state = "litellm-session-token:cli-test1234567890"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": "env_state_value"}):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert redirect_params["state"] == cli_state
assert code_verifier is None # PKCE not enabled by default
def test_state_priority_env_variable_when_no_cli_state(self):
"""
Test that GENERIC_CLIENT_STATE environment variable is used when CLI state is not provided
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
env_state = "custom_env_state_value"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": env_state}):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert redirect_params["state"] == env_state
assert code_verifier is None
def test_state_priority_generated_uuid_fallback(self):
"""
Test that a UUID is generated when neither CLI state nor env variable is provided
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange - no CLI state and no env variable
with patch.dict(os.environ, {}, clear=False):
# Remove GENERIC_CLIENT_STATE if it exists
os.environ.pop("GENERIC_CLIENT_STATE", None)
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert "state" in redirect_params
assert redirect_params["state"] is not None
assert len(redirect_params["state"]) == 32 # UUID hex is 32 chars
assert code_verifier is None
def test_state_with_pkce_enabled(self):
"""
Test that PKCE parameters are generated when GENERIC_CLIENT_USE_PKCE is enabled
"""
import base64
import hashlib
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
test_state = "test_state_123"
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert state
assert redirect_params["state"] == test_state
# Assert PKCE parameters
assert code_verifier is not None
assert len(code_verifier) == 43 # Standard PKCE verifier length
assert "code_challenge" in redirect_params
assert "code_challenge_method" in redirect_params
assert redirect_params["code_challenge_method"] == "S256"
# Verify code_challenge is correctly derived from code_verifier
expected_challenge_bytes = hashlib.sha256(
code_verifier.encode("utf-8")
).digest()
expected_challenge = (
base64.urlsafe_b64encode(expected_challenge_bytes)
.decode("utf-8")
.rstrip("=")
)
assert redirect_params["code_challenge"] == expected_challenge
def test_state_with_pkce_disabled(self):
"""
Test that PKCE parameters are NOT generated when GENERIC_CLIENT_USE_PKCE is false
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
test_state = "test_state_456"
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "false"}):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert redirect_params["state"] == test_state
assert code_verifier is None
assert "code_challenge" not in redirect_params
assert "code_challenge_method" not in redirect_params
def test_state_priority_cli_state_overrides_env_with_pkce(self):
"""
Test that CLI state takes priority over env variable even when PKCE is enabled
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
cli_state = "cli_state_priority"
env_state = "env_state_should_not_be_used"
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": env_state,
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=cli_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert redirect_params["state"] == cli_state # CLI state takes priority
assert redirect_params["state"] != env_state
# PKCE should still be generated
assert code_verifier is not None
assert "code_challenge" in redirect_params
assert "code_challenge_method" in redirect_params
def test_empty_string_state_uses_env_variable(self):
"""
Test that empty string state is treated as None and uses env variable
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
env_state = "env_state_for_empty_cli"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": env_state}):
# Act
(
redirect_params,
code_verifier,
) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state="", # Empty string
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert - empty string is falsy, so env variable should be used
# Note: This tests current implementation behavior
# If empty string should be treated differently, implementation needs update
assert redirect_params["state"] == env_state
def test_multiple_calls_generate_different_uuids(self):
"""
Test that multiple calls without state generate different UUIDs
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange - no state provided
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("GENERIC_CLIENT_STATE", None)
# Act
params1, _ = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
params2, _ = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=None,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
assert params1["state"] != params2["state"]
assert len(params1["state"]) == 32
assert len(params2["state"]) == 32
class TestPKCEFunctionality:
"""Test PKCE (Proof Key for Code Exchange) functionality"""
def test_generate_pkce_params(self):
"""
Test that generate_pkce_params generates valid PKCE parameters
"""
import base64
import hashlib
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Act
code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params()
# Assert
assert len(code_verifier) == 43
assert isinstance(code_verifier, str)
# Verify code_challenge is correctly generated from code_verifier
expected_challenge_bytes = hashlib.sha256(
code_verifier.encode("utf-8")
).digest()
expected_challenge = (
base64.urlsafe_b64encode(expected_challenge_bytes)
.decode("utf-8")
.rstrip("=")
)
assert code_challenge == expected_challenge
# Verify both are base64url encoded (no padding)
assert "=" not in code_verifier
assert "=" not in code_challenge
@pytest.mark.asyncio
async def test_prepare_token_exchange_parameters_with_pkce(self):
"""
Test prepare_token_exchange_parameters retrieves PKCE code_verifier from cache
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request with state parameter
mock_request = MagicMock(spec=Request)
test_state = "test_oauth_state_123"
mock_request.query_params = {"state": test_state}
# Mock cache with async methods — use dict format (primary path)
mock_cache = MagicMock(redis_cache=None)
test_code_verifier = "test_code_verifier_abc123xyz"
mock_cache.async_get_cache = AsyncMock(
return_value={"code_verifier": test_code_verifier}
)
mock_cache.async_delete_cache = AsyncMock()
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}),
):
# Act
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
)
# Assert
assert token_params["include_client_id"] is False
assert token_params["code_verifier"] == test_code_verifier
# Cache key is returned for deferred deletion (after exchange succeeds)
assert token_params["_pkce_cache_key"] == f"pkce_verifier:{test_state}"
# Verify cache was read but NOT deleted yet (deletion is deferred to after
# successful token exchange to preserve the verifier for retries)
mock_cache.async_get_cache.assert_called_once_with(
key=f"pkce_verifier:{test_state}"
)
mock_cache.async_delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_get_generic_sso_redirect_response_with_pkce(self):
"""
Test get_generic_sso_redirect_response with PKCE enabled stores verifier and adds challenge to URL
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock SSO provider
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
original_location = (
"https://auth.example.com/authorize?state=test456&client_id=abc"
)
mock_redirect_response.headers = {"location": original_location}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
mock_sso.__exit__ = MagicMock(return_value=False)
test_state = "test456"
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_set_cache = AsyncMock()
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
):
# Act
result = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state=test_state,
generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
# Verify async cache was called to store code_verifier
mock_cache.async_set_cache.assert_called_once()
cache_call = mock_cache.async_set_cache.call_args
assert cache_call.kwargs["key"] == f"pkce_verifier:{test_state}"
assert cache_call.kwargs["ttl"] == 600
# Value is stored as dict for proper JSON serialization in Redis
cached = cache_call.kwargs["value"]
assert isinstance(cached, dict) and "code_verifier" in cached
assert len(cached["code_verifier"]) == 43
# Verify PKCE parameters were added to the redirect URL
assert result is not None
updated_location = str(result.headers["location"])
assert "code_challenge=" in updated_location
assert "code_challenge_method=S256" in updated_location
assert f"state={test_state}" in updated_location
@pytest.mark.asyncio
async def test_pkce_redis_multi_pod_verifier_roundtrip(self):
"""
Mock Redis to verify PKCE code_verifier round-trip across "pods":
Pod A stores verifier in Redis; Pod B retrieves it (no real IdP).
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# In-memory mock of Redis (shared between "pods")
class MockRedisCache:
def __init__(self):
self._store = {}
async def async_set_cache(self, key, value, **kwargs):
self._store[key] = json.dumps(value)
async def async_get_cache(self, key, **kwargs):
val = self._store.get(key)
if val is None:
return None
# Simulate RedisCache._get_cache_logic: stored as JSON string, return decoded
if isinstance(val, str):
try:
return json.loads(val)
except (ValueError, TypeError):
return val
return val
async def async_delete_cache(self, key):
self._store.pop(key, None)
mock_redis = MockRedisCache()
mock_in_memory = MagicMock()
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
mock_redirect_response.headers = {
"location": "https://auth.example.com/authorize?state=multi_pod_state_xyz&client_id=abc"
}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
mock_sso.__exit__ = MagicMock(return_value=False)
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis):
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory
):
# Pod A: start login, store code_verifier in "Redis"
await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state="multi_pod_state_xyz",
generic_authorization_endpoint="https://auth.example.com/authorize",
)
mock_in_memory.async_set_cache.assert_not_called()
# MockRedisCache is a real class; assert on state, not .assert_called_*
stored_key = "pkce_verifier:multi_pod_state_xyz"
assert stored_key in mock_redis._store
stored_value = mock_redis._store[stored_key]
# Stored as JSON-serialized dict for Redis compatibility
stored_dict = json.loads(stored_value)
assert (
isinstance(stored_dict, dict) and "code_verifier" in stored_dict
)
assert len(stored_dict["code_verifier"]) == 43
# Pod B: callback with same state, retrieve from "Redis"
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "multi_pod_state_xyz"}
token_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert token_params["code_verifier"] == stored_dict["code_verifier"]
# Cache key returned for deferred deletion after successful exchange
assert token_params["_pkce_cache_key"] == stored_key
mock_in_memory.async_get_cache.assert_not_called()
# Deletion is deferred — key still present until exchange succeeds
assert stored_key in mock_redis._store
@pytest.mark.asyncio
async def test_pkce_fallback_in_memory_roundtrip_when_redis_none(self):
"""
Regression: When redis_usage_cache is None (no Redis configured),
code_verifier is stored and retrieved via user_api_key_cache.
Roundtrip works when callback hits same pod (same in-memory cache).
Single-pod or no-Redis deployments must continue to work.
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# In-memory store (simulates user_api_key_cache on one pod)
in_memory_store = {}
async def async_set_cache(key, value, **kwargs):
in_memory_store[key] = value
async def async_get_cache(key, **kwargs):
return in_memory_store.get(key)
async def async_delete_cache(key):
in_memory_store.pop(key, None)
mock_in_memory = MagicMock()
mock_in_memory.async_set_cache = AsyncMock(side_effect=async_set_cache)
mock_in_memory.async_get_cache = AsyncMock(side_effect=async_get_cache)
mock_in_memory.async_delete_cache = AsyncMock(side_effect=async_delete_cache)
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
mock_redirect_response.headers = {
"location": "https://auth.example.com/authorize?state=fallback_state_xyz&client_id=abc"
}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
mock_sso.__exit__ = MagicMock(return_value=False)
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.redis_usage_cache", None):
with patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory
):
# Pod A: start login, store code_verifier in in-memory cache
await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state="fallback_state_xyz",
generic_authorization_endpoint="https://auth.example.com/authorize",
)
mock_in_memory.async_set_cache.assert_called_once()
stored_key = mock_in_memory.async_set_cache.call_args.kwargs["key"]
stored_value = mock_in_memory.async_set_cache.call_args.kwargs[
"value"
]
assert stored_key == "pkce_verifier:fallback_state_xyz"
assert (
isinstance(stored_value, dict)
and len(stored_value["code_verifier"]) == 43
)
# Same pod: callback retrieves from in-memory cache
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "fallback_state_xyz"}
token_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert (
token_params["code_verifier"] == stored_value["code_verifier"]
)
# Cache key returned for deferred deletion after successful exchange
assert token_params["_pkce_cache_key"] == stored_key
mock_in_memory.async_get_cache.assert_called_once_with(
key=stored_key
)
# Deletion is deferred — not called by prepare_token_exchange_parameters
mock_in_memory.async_delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):
"""
Regression: prepare_token_exchange_parameters with no state in request
does not call cache and does not add code_verifier.
"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_redis = MagicMock()
mock_in_memory = MagicMock()
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory),
patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}, clear=False),
):
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
)
assert "code_verifier" not in token_params
mock_redis.async_get_cache.assert_not_called()
mock_in_memory.async_get_cache.assert_not_called()
@pytest.mark.asyncio
async def test_pkce_token_exchange_basic_auth(self):
"""When include_client_id=False, client credentials go via HTTP Basic Auth."""
token_resp = {
"access_token": "tok_abc",
"id_token": None,
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "user1", "email": "user@example.com"}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = token_resp
mock_userinfo_response = MagicMock()
mock_userinfo_response.status_code = 200
mock_userinfo_response.json.return_value = userinfo_resp
async def fake_post(*args, **kwargs):
# Verify Basic Auth is set via Authorization header
headers = kwargs.get("headers", {})
assert "Authorization" in headers
assert headers["Authorization"].startswith("Basic ")
# Verify code_verifier is in the POST body (essential PKCE field)
post_data = kwargs.get("data", {})
assert post_data.get("code_verifier") == "verifier_abc"
# Verify redirect_uri is forwarded (required by strict OAuth providers)
assert post_data.get("redirect_uri") == "https://proxy.example.com/callback"
# Verify credentials are NOT double-sent in the POST body when using Basic Auth
assert (
"client_secret" not in post_data
), "client_secret must not appear in POST body when using Basic Auth"
assert (
"client_id" not in post_data
), "client_id must not appear in POST body when using Basic Auth (include_client_id=False)"
return mock_response
# get_async_httpx_client returns a client directly (no context manager).
mock_token_client = MagicMock()
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = MagicMock()
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo_response)
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_get_client.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code_123",
code_verifier="verifier_abc",
client_id="my_client",
client_secret="my_secret",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_abc"
assert result["email"] == "user@example.com"
# id_token was explicit null in token_response — the merge loop must remove it
# rather than leaving "id_token": None in the result.
assert (
"id_token" not in result
), "null id_token from token endpoint must be absent in merged result"
# Verify userinfo GET used the correct Bearer token header
get_call = mock_userinfo_client.get.call_args
assert get_call is not None
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_abc"
@pytest.mark.asyncio
async def test_pkce_token_exchange_credentials_in_body(self):
"""When include_client_id=True, credentials go in the request body."""
token_resp = {
"access_token": "tok_body",
"id_token": None,
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "user2", "email": "user2@example.com"}
async def fake_post(*args, **kwargs):
headers = kwargs.get("headers", {})
auth_header = headers.get("Authorization", "")
assert not auth_header.startswith(
"Basic "
), "Should NOT use Basic Auth when include_client_id=True"
data = kwargs.get("data", {})
assert "client_id" in data
assert "client_secret" in data
assert (
data.get("code_verifier") == "verifier_xyz"
), "code_verifier must be in POST body"
assert (
data.get("redirect_uri") == "https://proxy.example.com/callback"
), "redirect_uri must be forwarded"
mock = MagicMock()
mock.status_code = 200
mock.json.return_value = token_resp
return mock
mock_userinfo = MagicMock()
mock_userinfo.status_code = 200
mock_userinfo.json.return_value = userinfo_resp
mock_token_client = MagicMock()
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = MagicMock()
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo)
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_get_client.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code_456",
code_verifier="verifier_xyz",
client_id="client_id_value",
client_secret="client_secret_value",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=True,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_body"
assert result["sub"] == "user2"
# Verify userinfo GET used the correct Bearer token header
get_call = mock_userinfo_client.get.call_args
assert get_call is not None
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_body"
@pytest.mark.asyncio
async def test_pkce_token_exchange_http200_with_error_body(self):
"""Provider returns HTTP 200 but with an error field instead of tokens."""
from litellm.proxy._types import ProxyException
error_body = {
"error": "invalid_grant",
"error_description": "Code already used",
}
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = error_body
mock_client.post = AsyncMock(return_value=mock_resp)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="expired_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert "invalid_grant" in exc_info.value.message
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_userinfo_falls_back_to_id_token(self):
"""When the userinfo endpoint fails, decode the id_token as fallback."""
import base64
import json as _json
payload = {"sub": "user_from_jwt", "email": "jwt@example.com"}
# Build a minimal JWT (header.payload.signature — signature not verified)
encoded_payload = (
base64.urlsafe_b64encode(_json.dumps(payload).encode())
.rstrip(b"=")
.decode()
)
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_fail = MagicMock()
mock_fail.status_code = 503
mock_client.get = AsyncMock(return_value=mock_fail)
mock_get_client.return_value = mock_client
result = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="some_token",
id_token=fake_id_token,
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert result["sub"] == "user_from_jwt"
assert result["email"] == "jwt@example.com"
@pytest.mark.asyncio
async def test_pkce_userinfo_uses_id_token_when_no_endpoint(self):
"""When userinfo_endpoint is None, fall back to id_token directly without HTTP call."""
import base64
import json as _json
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
payload = {"sub": "id_token_user", "email": "id@example.com"}
encoded_payload = (
base64.urlsafe_b64encode(_json.dumps(payload).encode())
.rstrip(b"=")
.decode()
)
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
# No httpx call should happen when userinfo_endpoint is None
result = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="some_token",
id_token=fake_id_token,
userinfo_endpoint=None,
additional_headers={},
)
assert result["sub"] == "id_token_user"
assert result["email"] == "id@example.com"
@pytest.mark.asyncio
async def test_pkce_userinfo_raises_when_both_sources_unavailable(self):
"""When userinfo endpoint fails AND no id_token, raise ProxyException."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_fail = MagicMock()
mock_fail.status_code = 503
mock_client.get = AsyncMock(return_value=mock_fail)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="token",
id_token=None, # no id_token available
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert "unavailable" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_userinfo_http200_empty_body_no_id_token_raises(self):
"""When userinfo returns HTTP 200 with an empty/null body and no id_token is
available, _get_pkce_userinfo raises ProxyException."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = None # HTTP 200 with null JSON body
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_client.get = AsyncMock(return_value=mock_resp)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="access_token",
id_token=None, # no id_token fallback available
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert (
"unavailable" in exc_info.value.message.lower()
or "no userinfo" in exc_info.value.message.lower()
or "userinfo" in exc_info.value.message.lower()
)
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_cache_miss_raises_proxy_exception(self):
"""prepare_token_exchange_parameters raises ProxyException when PKCE is enabled
but no verifier is found in cache (cross-instance cache miss scenario)."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "missing_state_123"}
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "true"},
),
):
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert (
"verifier not found" in exc_info.value.message.lower()
or "cache" in exc_info.value.message.lower()
)
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_token_exchange_public_client_no_secret(self):
"""Public PKCE client (include_client_id=False, no secret) sends client_id in
POST body and does NOT include Basic Auth or client_secret."""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
token_resp = {
"access_token": "tok_public",
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "pubuser", "email": "pub@example.com"}
async def fake_post(*args, **kwargs):
headers = kwargs.get("headers", {})
auth_header = headers.get("Authorization", "")
assert not auth_header.startswith(
"Basic "
), "Public client must not use Basic Auth"
data = kwargs.get("data", {})
assert data.get("client_id") == "public_client_id"
assert (
"client_secret" not in data
), "No secret should be sent for public client"
assert data.get("code_verifier") == "public_verifier"
mock = MagicMock()
mock.status_code = 200
mock.json.return_value = token_resp
return mock
mock_userinfo = MagicMock()
mock_userinfo.status_code = 200
mock_userinfo.json.return_value = userinfo_resp
mock_token_client = MagicMock()
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = MagicMock()
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo)
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_get_client.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_pub",
code_verifier="public_verifier",
client_id="public_client_id",
client_secret=None, # public client — no secret
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_public"
assert result["sub"] == "pubuser"
@pytest.mark.asyncio
async def test_delete_pkce_verifier_swallows_deletion_errors(self):
"""_delete_pkce_verifier must not raise when the cache delete fails
(best-effort cleanup — a leftover verifier must not abort a successful SSO login).
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
failing_cache = MagicMock()
failing_cache.async_delete_cache = AsyncMock(
side_effect=Exception("Redis down")
)
# Should NOT raise even though the underlying cache delete fails
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", failing_cache),
):
await SSOAuthenticationHandler._delete_pkce_verifier(
"pkce_verifier:test_state"
)
failing_cache.async_delete_cache.assert_called_once_with(
key="pkce_verifier:test_state"
)
@pytest.mark.asyncio
async def test_pkce_cache_miss_unexpected_format_raises_proxy_exception(self):
"""When cached data exists but has an unrecognized format (not a dict with
code_verifier, not a plain string), prepare_token_exchange_parameters raises
ProxyException rather than silently falling through to a non-PKCE flow."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Cache returns an integer — unexpected format
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_get_cache = AsyncMock(return_value=12345)
mock_cache.async_delete_cache = AsyncMock()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "bad_format_state"}
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "true"},
),
):
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert (
"cache" in exc_info.value.message.lower()
or "verifier" in exc_info.value.message.lower()
or "format" in exc_info.value.message.lower()
)
assert str(exc_info.value.code) == "401"
# Strict mode should also clean up the corrupt cache entry before raising
mock_cache.async_delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_pkce_cache_miss_non_strict_logs_warning_and_continues(self, caplog):
"""Default (non-strict) cache-miss behavior: logs a warning and returns params
without code_verifier rather than raising, to preserve backward compatibility.
"""
import logging
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "missing_state_non_strict"}
# PKCE_STRICT_CACHE_MISS explicitly set to false — should NOT raise.
# Use patch.dict with the key set to "false" rather than os.environ.pop()
# to avoid permanently mutating the test process environment.
with (
caplog.at_level(logging.WARNING),
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "false"},
clear=False,
),
):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
# Should return params without code_verifier (no raise)
assert "code_verifier" not in result
assert "_pkce_cache_key" not in result
# Non-strict mode emits a warning rather than raising
mock_cache.async_get_cache.assert_called_once()
# Verify the warning was actually logged
assert any(
"verifier not found" in r.message.lower()
or "code_verifier" in r.message.lower()
for r in caplog.records
if r.levelno >= logging.WARNING
), f"Expected a cache-miss warning. Records: {[r.message for r in caplog.records]}"
@pytest.mark.asyncio
async def test_pkce_token_exchange_non200_raises_proxy_exception(self):
"""_pkce_token_exchange raises ProxyException when the token endpoint
returns a non-200 status (e.g. 401 Unauthorized from provider)."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = "Unauthorized"
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code",
code_verifier="verifier",
client_id="client_id",
client_secret="secret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=True,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert "token" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_cache_miss_unexpected_format_non_strict_logs_warning(
self, caplog
):
"""When cached data has an unexpected format (e.g. integer from corrupt Redis)
in non-strict mode, prepare_token_exchange_parameters logs a warning and
returns params without code_verifier rather than raising."""
import logging
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Cache returns an integer — unexpected format
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_get_cache = AsyncMock(return_value=12345)
mock_cache.async_delete_cache = AsyncMock()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "bad_format_non_strict"}
# Non-strict mode: should log a warning and continue, not raise.
# Use patch.dict with PKCE_STRICT_CACHE_MISS="false" to avoid permanently
# mutating the test process environment with os.environ.pop().
with (
caplog.at_level(logging.WARNING),
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "false"},
clear=False,
),
):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
# No raise in non-strict mode; verifier simply absent from params
assert "code_verifier" not in result
assert "_pkce_cache_key" not in result
# Cache was queried (the unexpected format was retrieved and logged at WARNING)
mock_cache.async_get_cache.assert_called_once()
# Verify a warning was logged about the unexpected format or cache miss
assert any(
"verifier" in r.message.lower()
or "format" in r.message.lower()
or "cache" in r.message.lower()
for r in caplog.records
if r.levelno >= logging.WARNING
), f"Expected a format/cache warning. Records: {[r.message for r in caplog.records]}"
# Verify cleanup was attempted for the corrupt/stale cache entry
mock_cache.async_delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_pkce_legacy_string_cache_format_backward_compat(self):
"""Legacy plain-string cache entries (stored before dict format was introduced)
are handled transparently via the backward-compat branch."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
legacy_verifier = "legacy_plain_string_verifier_abc123"
mock_cache = MagicMock(redis_cache=None)
mock_cache.async_get_cache = AsyncMock(return_value=legacy_verifier)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "legacy_state_xyz"}
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}, clear=False),
):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert result["code_verifier"] == legacy_verifier
assert result["_pkce_cache_key"] == "pkce_verifier:legacy_state_xyz"
@pytest.mark.asyncio
async def test_pkce_token_exchange_null_json_body_raises_proxy_exception(self):
"""HTTP 200 with JSON body `null` raises a clean ProxyException instead of
AttributeError when .get() is called on the None return value."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = None # JSON null response body
mock_resp.text = "null"
mock_client.post = AsyncMock(return_value=mock_resp)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="some_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=False,
redirect_url=None,
additional_headers={},
)
assert "unexpected response format" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_token_exchange_http200_no_error_field_no_access_token(self):
"""HTTP 200 with no error field and no access_token raises ProxyException
with a descriptive message showing the actual response keys."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
body_without_token = {"token_type": "Bearer", "scope": "openid"}
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_get_client:
mock_client = MagicMock()
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = body_without_token
mock_client.post = AsyncMock(return_value=mock_resp)
mock_get_client.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="some_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=False,
redirect_url=None,
additional_headers={},
)
assert (
"no access_token" in exc_info.value.message
or "access_token" in exc_info.value.message
)
assert str(exc_info.value.code) == "401"
# Tests for SSO user team assignment bug (Issue: SSO Users Not Added to Entra-Synced Teams on First Login)
class TestAddMissingTeamMember:
"""Tests for the add_missing_team_member function"""
@pytest.mark.asyncio
async def test_add_missing_team_member_with_new_user_response_teams_none(self):
"""
Bug reproduction: When a NewUserResponse has teams=None (new SSO user),
add_missing_team_member() should still add the user to the SSO teams.
Currently FAILS: The function returns early when teams is None.
"""
from litellm.proxy._types import NewUserResponse
from litellm.proxy.management_endpoints.ui_sso import add_missing_team_member
# Simulate a new SSO user - NewUserResponse has teams=None by default
new_user = NewUserResponse(
user_id="new-sso-user-123",
key="sk-xxxxx",
teams=None, # This is the default for NewUserResponse
)
sso_teams = ["team-from-entra-1", "team-from-entra-2"]
with patch(
"litellm.proxy.management_endpoints.ui_sso.create_team_member_add_task"
) as mock_add_task:
mock_add_task.return_value = AsyncMock()
await add_missing_team_member(user_info=new_user, sso_teams=sso_teams)
# Bug: This assertion currently FAILS - no teams are added
# because function returns early when teams is None
assert (
mock_add_task.call_count == 2
), f"Expected 2 calls to add user to teams, but got {mock_add_task.call_count}"
called_team_ids = [call.args[0] for call in mock_add_task.call_args_list]
assert set(called_team_ids) == {
"team-from-entra-1",
"team-from-entra-2",
}
@pytest.mark.asyncio
async def test_add_missing_team_member_with_litellm_user_table_empty_teams(self):
"""
Control test: When a LiteLLM_UserTable has teams=[] (existing user, no teams),
add_missing_team_member() should add the user to SSO teams.
This test PASSES because LiteLLM_UserTable defaults teams to [] not None.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import add_missing_team_member
# Existing user has teams=[] by default (not None)
existing_user = LiteLLM_UserTable(
user_id="existing-user-456",
teams=[], # Empty list, not None
)
sso_teams = ["team-from-entra-1", "team-from-entra-2"]
with patch(
"litellm.proxy.management_endpoints.ui_sso.create_team_member_add_task"
) as mock_add_task:
mock_add_task.return_value = AsyncMock()
await add_missing_team_member(user_info=existing_user, sso_teams=sso_teams)
# This PASSES - teams are added because teams=[] not None
assert mock_add_task.call_count == 2
@pytest.mark.asyncio
async def test_add_user_to_teams_from_sso_response_new_user(self):
"""
Integration test: Simulates the SSO response handler with a new user
that has teams=None from NewUserResponse.
"""
from litellm.proxy._types import NewUserResponse
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
# SSO response with team_ids from Entra ID
sso_result = CustomOpenID(
id="new-sso-user-id",
email="newuser@example.com",
team_ids=["entra-group-1", "entra-group-2"],
)
# New user response (simulates what new_user() returns)
new_user_info = NewUserResponse(
user_id="new-sso-user-id",
key="sk-xxxxx",
teams=None, # Bug: NewUserResponse defaults to None
)
with patch(
"litellm.proxy.management_endpoints.ui_sso.add_missing_team_member"
) as mock_add_member:
await SSOAuthenticationHandler.add_user_to_teams_from_sso_response(
result=sso_result,
user_info=new_user_info,
)
# Verify add_missing_team_member was called with correct args
mock_add_member.assert_called_once_with(
user_info=new_user_info, sso_teams=["entra-group-1", "entra-group-2"]
)
@pytest.mark.asyncio
async def test_sso_first_login_full_flow_adds_user_to_teams(self):
"""
End-to-end test: Simulates complete first-time SSO login with Entra groups.
Verifies teams are created AND user is added as a member.
"""
from litellm.proxy._types import NewUserResponse
from litellm.proxy.management_endpoints.ui_sso import add_missing_team_member
team_member_calls = []
async def track_team_member_add(team_id, user_info):
team_member_calls.append({"team_id": team_id, "user_id": user_info.user_id})
# New SSO user with Entra groups
new_user = NewUserResponse(
user_id="first-time-sso-user",
key="sk-xxxxx",
teams=None, # The problematic default
)
sso_teams = ["entra-team-alpha", "entra-team-beta"]
with patch(
"litellm.proxy.management_endpoints.ui_sso.create_team_member_add_task",
side_effect=track_team_member_add,
):
await add_missing_team_member(user_info=new_user, sso_teams=sso_teams)
# Bug: With current code, team_member_calls will be empty
# After fix: Should have 2 entries
assert (
len(team_member_calls) == 2
), f"Expected 2 teams to be added, but got {len(team_member_calls)}"
assert {c["team_id"] for c in team_member_calls} == {
"entra-team-alpha",
"entra-team-beta",
}
assert all(c["user_id"] == "first-time-sso-user" for c in team_member_calls)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"user_info_factory,teams_value,expected_teams_added",
[
# Bug case: NewUserResponse with teams=None
pytest.param(
lambda uid: NewUserResponse(user_id=uid, key="sk-xxx", teams=None),
None,
["team-1", "team-2"], # Should still add teams
id="new_user_teams_none",
),
# Working case: LiteLLM_UserTable with teams=[]
pytest.param(
lambda uid: LiteLLM_UserTable(user_id=uid, teams=[]),
[],
["team-1", "team-2"],
id="existing_user_empty_teams",
),
# Existing user with some teams already
pytest.param(
lambda uid: LiteLLM_UserTable(user_id=uid, teams=["team-1"]),
["team-1"],
["team-2"], # Only missing team should be added
id="existing_user_partial_teams",
),
],
)
async def test_add_missing_team_member_handles_all_user_types(
self, user_info_factory, teams_value, expected_teams_added
):
"""
Parametrized test ensuring add_missing_team_member works for all user types.
"""
from litellm.proxy.management_endpoints.ui_sso import add_missing_team_member
user_info = user_info_factory("test-user-id")
sso_teams = ["team-1", "team-2"]
added_teams = []
async def mock_create_task(team_id, user):
added_teams.append(team_id)
with patch(
"litellm.proxy.management_endpoints.ui_sso.create_team_member_add_task",
side_effect=mock_create_task,
):
await add_missing_team_member(user_info=user_info, sso_teams=sso_teams)
assert set(added_teams) == set(
expected_teams_added
), f"Expected teams {expected_teams_added}, but got {added_teams}"
@pytest.mark.asyncio
async def test_role_mappings_override_default_internal_user_params():
"""
Test that when role_mappings is configured in SSO settings,
the SSO-extracted role overrides default_internal_user_params role.
"""
from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import insert_sso_user
# Save original default_internal_user_params
original_default_params = getattr(litellm, "default_internal_user_params", None)
try:
# Set default_internal_user_params with a role that should be overridden
litellm.default_internal_user_params = {
"user_role": "internal_user",
"max_budget": 100,
"budget_duration": "30d",
"models": ["gpt-3.5-turbo"],
}
# Mock SSO result
mock_result_openid = CustomOpenID(
id="test-user-123",
email="test@example.com",
display_name="Test User",
provider="microsoft",
team_ids=[],
)
# User defined values with SSO-extracted role (from role_mappings)
user_defined_values: SSOUserDefinedValues = {
"user_id": "test-user-123",
"user_email": "test@example.com",
"user_role": "proxy_admin", # Role from SSO role_mappings
"max_budget": None,
"budget_duration": None,
"models": [],
}
# Mock new_user function
mock_new_user_response = NewUserResponse(
user_id="test-user-123",
key="sk-xxxxx",
teams=None,
)
with patch(
"litellm.proxy.management_endpoints.ui_sso.new_user",
return_value=mock_new_user_response,
) as mock_new_user:
# Act
_ = await insert_sso_user(
result_openid=mock_result_openid,
user_defined_values=user_defined_values,
)
# Assert - verify new_user was called with preserved SSO role
mock_new_user.assert_called_once()
call_args = mock_new_user.call_args
new_user_request = call_args.kwargs["data"]
# The role from SSO should be preserved, not overridden by default_internal_user_params
assert (
new_user_request.user_role == "proxy_admin"
), "SSO-extracted role should override default_internal_user_params role"
# Other default params should still be applied
assert (
new_user_request.max_budget == 100
), "max_budget from default_internal_user_params should be applied"
assert (
new_user_request.budget_duration == "30d"
), "budget_duration from default_internal_user_params should be applied"
finally:
# Restore original default_internal_user_params (always assign, never delattr —
# the attribute is defined in litellm/__init__.py and delattr-ing it breaks parallel tests)
litellm.default_internal_user_params = original_default_params
@pytest.mark.asyncio
async def test_sso_role_preserved_without_role_mappings():
"""
Test that SSO-extracted role is preserved even when role_mappings is NOT configured.
This covers the case where the role comes from Microsoft app_roles or
GENERIC_USER_ROLE_ATTRIBUTE (not from LiteLLM's role_mappings feature).
Previously, the role was only preserved when role_mappings was configured,
causing admin users to be downgraded to internal_user.
"""
from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues
from litellm.proxy.management_endpoints.ui_sso import insert_sso_user
original_default_params = getattr(litellm, "default_internal_user_params", None)
try:
# Set default_internal_user_params (as most deployments do)
litellm.default_internal_user_params = {
"user_role": "internal_user",
"max_budget": 50,
}
# Mock SSO result from Microsoft with app_roles-derived admin role
mock_result_openid = CustomOpenID(
id="msft-user-456",
email="admin@company.com",
display_name="Admin User",
provider="microsoft",
team_ids=["group-1"],
user_role=None, # role is in user_defined_values, not on the OpenID result
)
# User defined values with role from Microsoft app_roles (NOT role_mappings)
user_defined_values: SSOUserDefinedValues = {
"user_id": "msft-user-456",
"user_email": "admin@company.com",
"user_role": "proxy_admin", # Role from Microsoft app_roles
"max_budget": None,
"budget_duration": None,
"models": [],
}
mock_new_user_response = NewUserResponse(
user_id="msft-user-456",
key="sk-xxxxx",
teams=None,
)
# No role_mappings configured anywhere - the role came from app_roles
with patch(
"litellm.proxy.management_endpoints.ui_sso.new_user",
return_value=mock_new_user_response,
) as mock_new_user:
_ = await insert_sso_user(
result_openid=mock_result_openid,
user_defined_values=user_defined_values,
)
mock_new_user.assert_called_once()
call_args = mock_new_user.call_args
new_user_request = call_args.kwargs["data"]
# SSO role should be preserved even without role_mappings configured
assert (
new_user_request.user_role == "proxy_admin"
), "SSO role from app_roles should not be overwritten by default_internal_user_params"
# Other defaults should still apply
assert (
new_user_request.max_budget == 50
), "max_budget from default_internal_user_params should be applied"
# Verify user_defined_values was also updated (it's mutated in-place)
assert (
user_defined_values["user_role"] == "proxy_admin"
), "user_defined_values should retain the SSO role after insert_sso_user"
finally:
# Restore original default_internal_user_params (always assign, never delattr —
# deleting the attribute causes AttributeError in subsequent tests because
# litellm.__getattr__ has no handler for this name)
litellm.default_internal_user_params = original_default_params
class TestSSOReadinessEndpoint:
"""Test the /sso/readiness endpoint"""
@pytest.mark.asyncio
async def test_sso_readiness_no_sso_configured(self):
"""Test that readiness returns healthy when no SSO is configured"""
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
client = TestClient(app)
with patch.dict(os.environ, {}, clear=True):
response = client.get("/sso/readiness")
assert response.status_code == 200
data = response.json()
assert data["status"] == "healthy"
assert data["sso_configured"] is False
assert data["message"] == "No SSO provider configured"
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_sso_readiness_google_fully_configured(self):
"""Test that readiness returns healthy when Google SSO is fully configured"""
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
client = TestClient(app)
with patch.dict(
os.environ,
{
"GOOGLE_CLIENT_ID": "test-google-client-id",
"GOOGLE_CLIENT_SECRET": "test-google-secret",
},
clear=True,
):
response = client.get("/sso/readiness")
assert response.status_code == 200
data = response.json()
assert data["status"] == "healthy"
assert data["sso_configured"] is True
assert data["provider"] == "google"
assert "Google SSO is properly configured" in data["message"]
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_sso_readiness_google_missing_secret(self):
"""Test that readiness returns unhealthy when Google SSO is missing GOOGLE_CLIENT_SECRET"""
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
client = TestClient(app)
with patch.dict(
os.environ,
{"GOOGLE_CLIENT_ID": "test-google-client-id"},
clear=True,
):
response = client.get("/sso/readiness")
assert response.status_code == 503
data = response.json()["detail"]
assert data["status"] == "unhealthy"
assert data["sso_configured"] is True
assert data["provider"] == "google"
assert "GOOGLE_CLIENT_SECRET" in data["missing_environment_variables"]
assert (
"Google SSO is configured but missing required environment variables"
in data["message"]
)
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"env_vars,expected_status,expected_provider,expected_missing_vars",
[
(
{
"MICROSOFT_CLIENT_ID": "test-microsoft-client-id",
"MICROSOFT_CLIENT_SECRET": "test-microsoft-secret",
"MICROSOFT_TENANT": "test-tenant",
},
200,
"microsoft",
[],
),
(
{"MICROSOFT_CLIENT_ID": "test-microsoft-client-id"},
503,
"microsoft",
["MICROSOFT_CLIENT_SECRET", "MICROSOFT_TENANT"],
),
],
)
async def test_sso_readiness_microsoft_configurations(
self, env_vars, expected_status, expected_provider, expected_missing_vars
):
"""Test Microsoft SSO readiness with both fully configured and missing variables"""
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
client = TestClient(app)
with patch.dict(os.environ, env_vars, clear=True):
response = client.get("/sso/readiness")
assert response.status_code == expected_status
if expected_status == 200:
data = response.json()
assert data["sso_configured"] is True
assert data["provider"] == expected_provider
assert data["status"] == "healthy"
assert "Microsoft SSO is properly configured" in data["message"]
else:
data = response.json()["detail"]
assert data["sso_configured"] is True
assert data["provider"] == expected_provider
assert data["status"] == "unhealthy"
assert set(data["missing_environment_variables"]) == set(
expected_missing_vars
)
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"env_vars,expected_status,expected_provider,expected_missing_vars",
[
(
{
"GENERIC_CLIENT_ID": "test-generic-client-id",
"GENERIC_CLIENT_SECRET": "test-generic-secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/authorize",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
},
200,
"generic",
[],
),
(
{"GENERIC_CLIENT_ID": "test-generic-client-id"},
503,
"generic",
[
"GENERIC_CLIENT_SECRET",
"GENERIC_AUTHORIZATION_ENDPOINT",
"GENERIC_TOKEN_ENDPOINT",
"GENERIC_USERINFO_ENDPOINT",
],
),
],
)
async def test_sso_readiness_generic_configurations(
self, env_vars, expected_status, expected_provider, expected_missing_vars
):
"""Test Generic SSO readiness with both fully configured and missing variables"""
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import app
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
try:
client = TestClient(app)
with patch.dict(os.environ, env_vars, clear=True):
response = client.get("/sso/readiness")
assert response.status_code == expected_status
if expected_status == 200:
data = response.json()
assert data["sso_configured"] is True
assert data["provider"] == expected_provider
assert data["status"] == "healthy"
assert "Generic SSO is properly configured" in data["message"]
else:
data = response.json()["detail"]
assert data["sso_configured"] is True
assert data["provider"] == expected_provider
assert data["status"] == "unhealthy"
assert set(data["missing_environment_variables"]) == set(
expected_missing_vars
)
finally:
app.dependency_overrides.clear()
class TestCustomMicrosoftSSO:
"""Tests for CustomMicrosoftSSO class."""
@pytest.mark.asyncio
async def test_custom_microsoft_sso_uses_default_endpoints_when_no_env_vars(self):
"""
Test that CustomMicrosoftSSO uses default Microsoft endpoints
when no custom environment variables are set.
"""
# Ensure no custom endpoints are set
for key in [
"MICROSOFT_AUTHORIZATION_ENDPOINT",
"MICROSOFT_TOKEN_ENDPOINT",
"MICROSOFT_USERINFO_ENDPOINT",
]:
os.environ.pop(key, None)
sso = CustomMicrosoftSSO(
client_id="test-client-id",
client_secret="test-client-secret",
tenant="test-tenant",
redirect_uri="http://localhost:4000/sso/callback",
)
discovery = await sso.get_discovery_document()
assert (
discovery["authorization_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/authorize"
)
assert (
discovery["token_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
)
assert discovery["userinfo_endpoint"] == "https://graph.microsoft.com/v1.0/me"
@pytest.mark.asyncio
async def test_custom_microsoft_sso_uses_custom_endpoints_when_env_vars_set(self):
"""
Test that CustomMicrosoftSSO uses custom endpoints
when environment variables are set.
"""
custom_auth_endpoint = "https://custom.example.com/oauth2/v2.0/authorize"
custom_token_endpoint = "https://custom.example.com/oauth2/v2.0/token"
custom_userinfo_endpoint = "https://custom.example.com/v1.0/me"
with patch.dict(
os.environ,
{
"MICROSOFT_AUTHORIZATION_ENDPOINT": custom_auth_endpoint,
"MICROSOFT_TOKEN_ENDPOINT": custom_token_endpoint,
"MICROSOFT_USERINFO_ENDPOINT": custom_userinfo_endpoint,
},
):
sso = CustomMicrosoftSSO(
client_id="test-client-id",
client_secret="test-client-secret",
tenant="test-tenant",
redirect_uri="http://localhost:4000/sso/callback",
)
discovery = await sso.get_discovery_document()
assert discovery["authorization_endpoint"] == custom_auth_endpoint
assert discovery["token_endpoint"] == custom_token_endpoint
assert discovery["userinfo_endpoint"] == custom_userinfo_endpoint
@pytest.mark.asyncio
async def test_custom_microsoft_sso_uses_partial_custom_endpoints(self):
"""
Test that CustomMicrosoftSSO uses custom endpoints for those set,
and defaults for others.
"""
custom_auth_endpoint = "https://custom.example.com/oauth2/v2.0/authorize"
# Clear other env vars first
os.environ.pop("MICROSOFT_TOKEN_ENDPOINT", None)
os.environ.pop("MICROSOFT_USERINFO_ENDPOINT", None)
with patch.dict(
os.environ,
{
"MICROSOFT_AUTHORIZATION_ENDPOINT": custom_auth_endpoint,
},
):
sso = CustomMicrosoftSSO(
client_id="test-client-id",
client_secret="test-client-secret",
tenant="test-tenant",
redirect_uri="http://localhost:4000/sso/callback",
)
discovery = await sso.get_discovery_document()
# Custom auth endpoint
assert discovery["authorization_endpoint"] == custom_auth_endpoint
# Default token and userinfo endpoints
assert (
discovery["token_endpoint"]
== "https://login.microsoftonline.com/test-tenant/oauth2/v2.0/token"
)
assert (
discovery["userinfo_endpoint"] == "https://graph.microsoft.com/v1.0/me"
)
def test_custom_microsoft_sso_uses_common_tenant_when_none(self):
"""
Test that CustomMicrosoftSSO uses 'common' tenant when tenant is None.
"""
sso = CustomMicrosoftSSO(
client_id="test-client-id",
client_secret="test-client-secret",
tenant=None,
redirect_uri="http://localhost:4000/sso/callback",
)
assert sso.tenant == "common"
def test_custom_microsoft_sso_is_subclass_of_microsoft_sso(self):
"""
Test that CustomMicrosoftSSO is a subclass of MicrosoftSSO.
"""
from fastapi_sso.sso.microsoft import MicrosoftSSO
sso = CustomMicrosoftSSO(
client_id="test-client-id",
client_secret="test-client-secret",
tenant="test-tenant",
redirect_uri="http://localhost:4000/sso/callback",
)
assert isinstance(sso, MicrosoftSSO)
@pytest.mark.asyncio
async def test_setup_team_mappings():
"""Test _setup_team_mappings function loads team mappings from database."""
# Arrange
mock_prisma = MagicMock()
mock_sso_config = MagicMock()
mock_sso_config.sso_settings = {"team_mappings": {"team_ids_jwt_field": "groups"}}
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(
return_value=mock_sso_config
)
with patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=mock_prisma,
):
# Act
result = await _setup_team_mappings()
# Assert
assert result is not None
assert isinstance(result, TeamMappings)
assert result.team_ids_jwt_field == "groups"
mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once_with(
where={"id": "sso_config"}
)
# ============================================================================
# Tests for get_litellm_user_role with list inputs (Keycloak returns lists)
# ============================================================================
def test_get_litellm_user_role_with_string():
"""Test that get_litellm_user_role works with a plain string."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
result = get_litellm_user_role("proxy_admin")
assert result == LitellmUserRoles.PROXY_ADMIN
def test_get_litellm_user_role_with_list():
"""
Test that get_litellm_user_role handles list inputs.
Keycloak returns roles as arrays like ["proxy_admin"] instead of strings.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
result = get_litellm_user_role(["proxy_admin"])
assert result == LitellmUserRoles.PROXY_ADMIN
def test_get_litellm_user_role_with_empty_list():
"""Test that get_litellm_user_role returns None for empty lists."""
from litellm.proxy.management_endpoints.types import get_litellm_user_role
result = get_litellm_user_role([])
assert result is None
def test_get_litellm_user_role_with_invalid_role():
"""Test that get_litellm_user_role returns None for invalid roles."""
from litellm.proxy.management_endpoints.types import get_litellm_user_role
result = get_litellm_user_role("not_a_real_role")
assert result is None
@pytest.mark.parametrize(
"role_claim",
[
["proxy_admin", "internal_user"],
["internal_user", "proxy_admin"],
],
)
def test_get_litellm_user_role_picks_highest_privilege_regardless_of_order(role_claim):
"""A multi-valued role claim resolves to the most privileged role, not the first one listed."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["proxy_admin_viewer", "internal_user"],
["internal_user", "proxy_admin_viewer"],
],
)
def test_get_litellm_user_role_keeps_org_spend_visibility_for_mixed_roles(role_claim):
"""
Regression for LIT-6077: a user holding both proxy_admin_viewer and internal_user kept
losing org-level spend visibility whenever the IdP happened to list internal_user first.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
def test_get_litellm_user_role_ignores_unrecognised_entries():
"""Roles LiteLLM does not know about are skipped rather than swallowing the whole claim."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(["some_idp_group", "internal_user"]) == LitellmUserRoles.INTERNAL_USER
assert get_litellm_user_role(["some_idp_group", "another_group"]) is None
def test_get_litellm_user_role_list_lookup_is_case_insensitive():
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(["INTERNAL_USER", "Proxy_Admin"]) == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["org_admin", "team"],
["team", "org_admin"],
],
)
def test_get_litellm_user_role_is_deterministic_for_unranked_roles(role_claim):
"""Roles outside the privilege hierarchy still resolve the same way in either claim order."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.ORG_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["org_admin", "internal_user"],
["internal_user", "org_admin"],
],
)
def test_get_litellm_user_role_prefers_a_ranked_role_over_an_unranked_one(role_claim):
"""
org_admin, team and customer sit outside the privilege ladder, so a claim mixing one of
them with a ranked role settles on the ranked role in either order. Same rule the Entra
app_roles and role_mappings paths already follow.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.INTERNAL_USER
def test_get_litellm_user_role_returns_none_for_non_string_claims():
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(None) is None
assert get_litellm_user_role({"role": "proxy_admin"}) is None
# ============================================================================
# Tests for process_sso_jwt_access_token role extraction
# ============================================================================
def test_process_sso_jwt_access_token_extracts_role_from_access_token():
"""
Test that process_sso_jwt_access_token extracts user role from the JWT
access token when the UserInfo response did not include it.
This is the core fix for the Keycloak SSO role mapping bug: Keycloak's
UserInfo endpoint does not return role claims, but the JWT access token
contains them.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
# Create a JWT access token with role claims (as Keycloak would)
access_token_payload = {
"sub": "user-123",
"email": "admin@test.com",
"litellm_role": ["proxy_admin"],
}
access_token_str = pyjwt.encode(access_token_payload, "secret", algorithm="HS256")
# Result object with no role set (simulating UserInfo response without roles)
result = CustomOpenID(
id="user-123",
email="admin@test.com",
display_name="Admin User",
team_ids=[],
user_role=None,
)
# Call with GENERIC_USER_ROLE_ATTRIBUTE pointing to litellm_role
with patch.dict(os.environ, {"GENERIC_USER_ROLE_ATTRIBUTE": "litellm_role"}):
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=None,
)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["internal_user", "proxy_admin_viewer"],
["proxy_admin_viewer", "internal_user"],
],
)
def test_process_sso_jwt_access_token_resolves_highest_privilege_role(role_claim):
"""
The generic SSO access-token path must land on the same role for a user whose role
claim holds several roles, whichever order the IdP emitted them in.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
access_token_str = pyjwt.encode(
{"sub": "user-123", "email": "mixed@test.com", "litellm_role": role_claim},
"secret",
algorithm="HS256",
)
result = CustomOpenID(
id="user-123",
email="mixed@test.com",
display_name="Mixed Role User",
team_ids=[],
user_role=None,
)
with patch.dict(os.environ, {"GENERIC_USER_ROLE_ATTRIBUTE": "litellm_role"}):
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=None,
)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
def test_process_sso_jwt_access_token_does_not_override_existing_role():
"""
Test that process_sso_jwt_access_token does NOT override a role that was
already extracted from the UserInfo response.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
access_token_payload = {
"sub": "user-123",
"litellm_role": ["internal_user"],
}
access_token_str = pyjwt.encode(access_token_payload, "secret", algorithm="HS256")
# Result already has a role (e.g., set from UserInfo)
result = CustomOpenID(
id="user-123",
email="admin@test.com",
display_name="Admin User",
team_ids=[],
user_role=LitellmUserRoles.PROXY_ADMIN,
)
with patch.dict(os.environ, {"GENERIC_USER_ROLE_ATTRIBUTE": "litellm_role"}):
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=None,
)
# Should keep the original role
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_process_sso_jwt_access_token_extracts_role_from_nested_field():
"""
Test role extraction from a nested JWT field like resource_access.client.roles.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
access_token_payload = {
"sub": "user-123",
"resource_access": {"my-client": {"roles": ["proxy_admin"]}},
}
access_token_str = pyjwt.encode(access_token_payload, "secret", algorithm="HS256")
result = CustomOpenID(
id="user-123",
email="admin@test.com",
display_name="Admin User",
team_ids=[],
user_role=None,
)
with patch.dict(
os.environ, {"GENERIC_USER_ROLE_ATTRIBUTE": "resource_access.my-client.roles"}
):
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=None,
)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_process_sso_jwt_access_token_with_role_mappings():
"""
Test role extraction using role_mappings (group-based role determination)
from the JWT access token.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings
access_token_payload = {
"sub": "user-123",
"groups": ["keycloak-admins", "developers"],
}
access_token_str = pyjwt.encode(access_token_payload, "secret", algorithm="HS256")
result = CustomOpenID(
id="user-123",
email="admin@test.com",
display_name="Admin User",
team_ids=[],
user_role=None,
)
role_mappings = RoleMappings(
provider="generic",
group_claim="groups",
default_role=LitellmUserRoles.INTERNAL_USER,
roles={
LitellmUserRoles.PROXY_ADMIN: ["keycloak-admins"],
LitellmUserRoles.INTERNAL_USER: ["developers"],
},
)
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=role_mappings,
)
# Should get highest privilege role
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_generic_response_convertor_with_extra_attributes(monkeypatch):
"""Test that extra attributes are extracted when GENERIC_USER_EXTRA_ATTRIBUTES is set"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
monkeypatch.setenv("GENERIC_CLIENT_ID", "test_client")
monkeypatch.setenv(
"GENERIC_USER_EXTRA_ATTRIBUTES", "custom_field1,custom_field2,custom_field3"
)
mock_response = {
"sub": "user-id-123",
"email": "user@example.com",
"given_name": "John",
"family_name": "Doe",
"name": "John Doe",
"provider": "generic",
"custom_field1": "value1",
"custom_field2": ["item1", "item2"],
"custom_field3": {"nested": "data"},
}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = generic_response_convertor(
response=mock_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
role_mappings=None,
)
assert result.extra_fields is not None
assert result.extra_fields["custom_field1"] == "value1"
assert result.extra_fields["custom_field2"] == ["item1", "item2"]
assert result.extra_fields["custom_field3"] == {"nested": "data"}
def test_generic_response_convertor_without_extra_attributes(monkeypatch):
"""Test backward compatibility - extra_fields is None when env var not set"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
monkeypatch.setenv("GENERIC_CLIENT_ID", "test_client")
# Don't set GENERIC_USER_EXTRA_ATTRIBUTES
mock_response = {
"sub": "user-id-123",
"email": "user@example.com",
"given_name": "John",
"family_name": "Doe",
"name": "John Doe",
"provider": "generic",
"custom_field1": "value1",
"custom_field2": "value2",
}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = generic_response_convertor(
response=mock_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
role_mappings=None,
)
assert result.extra_fields is None
def test_generic_response_convertor_extra_attributes_with_nested_paths(monkeypatch):
"""Test that nested paths work with dot notation"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
monkeypatch.setenv("GENERIC_CLIENT_ID", "test_client")
monkeypatch.setenv(
"GENERIC_USER_EXTRA_ATTRIBUTES", "org_info.department,org_info.manager"
)
mock_response = {
"sub": "user-id-123",
"email": "user@example.com",
"org_info": {"department": "Engineering", "manager": "Jane Smith"},
}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = generic_response_convertor(
response=mock_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
role_mappings=None,
)
assert result.extra_fields is not None
assert result.extra_fields["org_info.department"] == "Engineering"
assert result.extra_fields["org_info.manager"] == "Jane Smith"
def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch):
"""Test that missing fields return None"""
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
monkeypatch.setenv("GENERIC_CLIENT_ID", "test_client")
monkeypatch.setenv("GENERIC_USER_EXTRA_ATTRIBUTES", "missing_field,another_missing")
mock_response = {
"sub": "user-id-123",
"email": "user@example.com",
}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
result = generic_response_convertor(
response=mock_response,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
role_mappings=None,
)
assert result.extra_fields is not None
assert result.extra_fields["missing_field"] is None
assert result.extra_fields["another_missing"] is None
class TestCliSsoAttributionMetadata:
"""CLI SSO allowlisted OIDC claim persistence and poll exposure."""
def test_parse_cli_sso_claim_map(self, monkeypatch):
from litellm.proxy.management_endpoints import ui_sso
monkeypatch.setattr(
ui_sso,
"CLI_SSO_CLAIM_MAP",
"employment_type->metadata.acme_employment_type, org_info.department -> department",
)
assert ui_sso._parse_cli_sso_claim_map() == [
("employment_type", "acme_employment_type"),
("org_info.department", "department"),
]
def test_build_cli_sso_attribution_metadata_filters_non_scalars(self, monkeypatch):
from litellm.proxy.management_endpoints import ui_sso
from litellm.proxy.management_endpoints.types import CustomOpenID
monkeypatch.setattr(
ui_sso,
"CLI_SSO_CLAIM_MAP",
"employment_type->acme_employment_type,access_token->should_drop,group->groups",
)
result = CustomOpenID(
id="user-1",
email="user@example.com",
display_name="User",
provider="generic",
team_ids=[],
extra_fields={
"employment_type": "full_time",
"access_token": "eyJhbGciOiJIUzI1NiJ9.payload.signature",
"group": ["team-a", "team-b"],
},
)
metadata = ui_sso.build_cli_sso_attribution_metadata(result=result)
assert metadata == {"acme_employment_type": "full_time"}
def test_build_cli_sso_attribution_metadata_from_oidc_dict(self, monkeypatch):
from litellm.proxy.management_endpoints import ui_sso
monkeypatch.setattr(
ui_sso,
"CLI_SSO_CLAIM_MAP",
"org_info.department->department",
)
metadata = ui_sso.build_cli_sso_attribution_metadata(
result={
"sub": "user-1",
"email": "user@example.com",
"org_info": {"department": "Engineering"},
}
)
assert metadata == {"department": "Engineering"}
@pytest.mark.asyncio
async def test_cli_sso_callback_passes_user_defined_values_for_new_users(self):
"""First CLI SSO login must supply SSOUserDefinedValues so upsert can create the user."""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints import ui_sso
from litellm.proxy.management_endpoints.types import CustomOpenID
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
session_key = "cli-session-new-user"
mock_user_info = LiteLLM_UserTable(
user_id="cli-test-user",
user_role="internal_user",
teams=[],
models=[],
)
mock_sso_result = CustomOpenID(
id="cli-test-user",
email="cli-test@example.com",
display_name="cli-test-user",
provider="generic",
team_ids=[],
)
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
get_user_info_mock = AsyncMock(return_value=mock_user_info)
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
get_user_info_mock,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.user_custom_sso", None),
):
await ui_sso.cli_sso_callback(
request=mock_request,
key=session_key,
result=mock_sso_result,
)
get_user_info_mock.assert_awaited_once()
assert get_user_info_mock.call_args.kwargs["user_defined_values"] is not None
assert (
get_user_info_mock.call_args.kwargs["user_defined_values"]["user_id"]
== "cli-test-user"
)
@pytest.mark.asyncio
async def test_cli_sso_callback_rejects_restricted_sso_group(self):
"""CLI SSO must enforce restricted_sso_group before upserting the user."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints import ui_sso
from litellm.proxy.management_endpoints.types import CustomOpenID
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://internal-proxy.local/"
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
mock_sso_result = CustomOpenID(
id="cli-test-user",
email="cli-test@example.com",
display_name="cli-test-user",
provider="generic",
team_ids=["other-group"],
)
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
new=AsyncMock(),
) as get_user_info_mock,
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.user_custom_sso", None),
patch(
"litellm.proxy.proxy_server.general_settings",
{
"ui_access_mode": {
"type": "restricted_sso_group",
"restricted_sso_group": "required-group",
}
},
),
):
with pytest.raises(ProxyException):
await ui_sso.cli_sso_callback(
request=mock_request,
key="cli-session-restricted",
result=mock_sso_result,
received_response={"groups": ["other-group"]},
)
get_user_info_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_cli_sso_callback_persists_attribution_metadata(self, monkeypatch):
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints import ui_sso
monkeypatch.setattr(
ui_sso,
"CLI_SSO_CLAIM_MAP",
"employment_type->acme_employment_type",
)
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
session_key = "cli-session-4567890"
mock_user_info = LiteLLM_UserTable(
user_id="test-user-123",
user_role="internal_user",
teams=["team1"],
models=["gpt-4"],
)
mock_sso_result = {
"user_email": "test@example.com",
"user_id": "test-user-123",
"employment_type": "contractor",
}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(
return_value=MagicMock(metadata={"auth_provider": "generic"})
)
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(
return_value=[
MagicMock(
model_dump=lambda: {
"team_id": "team1",
"team_alias": "team1",
"models": [],
}
)
]
)
with (
patch.dict(
os.environ,
{
"PROXY_BASE_URL": "https://test.example.com",
"SERVER_ROOT_PATH": "",
},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
return_value=mock_user_info,
),
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.user_custom_sso", None),
patch(
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
return_value="<html>Success</html>",
),
):
await ui_sso.cli_sso_callback(
request=mock_request,
key=session_key,
result=mock_sso_result,
)
flow_data = mock_cache.set_cache.call_args.kwargs["value"]
assert flow_data["session_data"]["attribution_metadata"] == {
"acme_employment_type": "contractor"
}
mock_prisma.db.litellm_usertable.update_many.assert_awaited_once()
update_data = mock_prisma.db.litellm_usertable.update_many.call_args.kwargs[
"data"
]
assert update_data["metadata"]["acme_employment_type"] == "contractor"
assert update_data["metadata"]["auth_provider"] == "generic"
@pytest.mark.asyncio
async def test_cli_poll_key_returns_attribution_metadata(self, monkeypatch):
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_key = "cli-session-789123"
session_data = {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": ["team-a", "team-b"],
"models": ["gpt-4"],
"attribution_metadata": {
"acme_employment_type": "full_time",
"org": {"cost_center": "CC-42"},
},
}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
):
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["attribution_metadata"] == {
"acme_employment_type": "full_time",
"org.cost_center": "CC-42",
}
class TestValidateReturnTo:
"""Tests for SSOAuthenticationHandler._validate_return_to"""
def test_returns_false_when_no_control_plane_url_configured(self, monkeypatch):
"""return_to should be silently ignored if control_plane_url is not in general_settings."""
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
result = SSOAuthenticationHandler._validate_return_to(
"https://cp.example.com/ui"
)
assert result is False
def test_allows_matching_origin(self, monkeypatch):
"""return_to matching the configured control_plane_url origin should pass."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
# Should not raise
SSOAuthenticationHandler._validate_return_to(
"https://cp.example.com/ui?page=models"
)
def test_allows_matching_origin_with_trailing_slash(self, monkeypatch):
"""Trailing slash on control_plane_url should not affect origin comparison."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com/"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
def test_rejects_prefix_attack(self, monkeypatch):
"""return_to like cp.example.com.evil.com must be rejected (not just prefix match)."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to(
"https://cp.example.com.evil.com/steal"
)
assert exc_info.value.status_code == 400
def test_rejects_different_origin(self, monkeypatch):
"""return_to pointing to a completely different domain should be rejected."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("https://evil.com/phish")
assert exc_info.value.status_code == 400
def test_case_insensitive_hostname(self, monkeypatch):
"""Hostname comparison should be case-insensitive per RFC 3986."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://CP.Example.COM"},
)
# Should not raise
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
def test_rejects_scheme_mismatch(self, monkeypatch):
"""http:// must be rejected when control_plane_url uses https://."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("http://cp.example.com/ui")
assert exc_info.value.status_code == 400
def test_rejects_port_mismatch(self, monkeypatch):
"""Non-default port must be rejected."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to(
"https://cp.example.com:8443/ui"
)
assert exc_info.value.status_code == 400
def test_allows_explicit_default_port(self, monkeypatch):
"""https://host:443 should match https://host (default port normalisation)."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:443/ui")
def test_allows_matching_custom_port(self, monkeypatch):
"""Both sides on the same custom port should match."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com:3000"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui")
class TestSyncUserRoleFromJwtRoleMap:
"""Tests for _sync_user_role_from_jwt_role_map."""
@staticmethod
def _make_jwt_handler():
from litellm.caching.caching import DualCache
from litellm.proxy._types import (
JWTLiteLLMRoleMap,
LiteLLM_JWTAuth,
LitellmUserRoles,
)
handler = JWTHandler()
handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(
roles_jwt_field="custom_roles",
user_id_upsert=True,
sync_user_role_and_teams=True,
jwt_litellm_role_map=[
JWTLiteLLMRoleMap(
jwt_role="my-admin",
litellm_role=LitellmUserRoles.PROXY_ADMIN,
),
JWTLiteLLMRoleMap(
jwt_role="my-viewer",
litellm_role=LitellmUserRoles.INTERNAL_USER,
),
],
),
)
return handler
@staticmethod
def _make_sso_values(user_role=None):
from litellm.proxy._types import SSOUserDefinedValues
user_id = "testuser@example.com"
return SSOUserDefinedValues(
models=[],
user_id=user_id,
user_email=user_id,
user_role=user_role,
max_budget=None,
budget_duration=None,
)
@pytest.mark.asyncio
async def test_stripped_response_has_no_roles(self):
"""Bug repro: stripped received_response lacks role claims."""
from litellm.caching.caching import DualCache
handler = self._make_jwt_handler()
sso_values = self._make_sso_values()
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"token_type": "Bearer", "expires_in": 3600},
user_info=None,
prisma_client=AsyncMock(),
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
assert sso_values["user_role"] is None
@pytest.mark.asyncio
async def test_decoded_access_token_maps_role(self):
"""Decoded JWT payload with role claims maps correctly."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
sso_values = self._make_sso_values()
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={
"sub": "testuser@example.com",
"custom_roles": ["my-admin"],
},
user_info=None,
prisma_client=AsyncMock(),
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
@pytest.mark.asyncio
async def test_existing_user_role_updated_in_db_and_cache(self):
"""Existing user with stale role gets updated in DB and cache."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
cache = DualCache()
prisma = AsyncMock()
prisma.db.litellm_usertable.update = AsyncMock()
user_id = "testuser@example.com"
existing_user = LiteLLM_UserTable(
user_id=user_id,
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
)
await cache.async_set_cache(
key=user_id, value=existing_user.model_dump(), ttl=60
)
sso_values = self._make_sso_values(
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
)
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"sub": user_id, "custom_roles": ["my-admin"]},
user_info=existing_user,
prisma_client=prisma,
user_api_key_cache=cache,
user_defined_values=sso_values,
)
prisma.db.litellm_usertable.update.assert_called_once_with(
where={"user_id": user_id},
data={"user_role": LitellmUserRoles.PROXY_ADMIN.value},
)
assert existing_user.user_role == LitellmUserRoles.PROXY_ADMIN.value
assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
@pytest.mark.asyncio
async def test_same_role_no_db_write(self):
"""No DB update when the mapped role matches the existing role."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
prisma = AsyncMock()
prisma.db.litellm_usertable.update = AsyncMock()
existing_user = LiteLLM_UserTable(
user_id="testuser@example.com",
user_role=LitellmUserRoles.PROXY_ADMIN.value,
)
sso_values = self._make_sso_values(
user_role=LitellmUserRoles.PROXY_ADMIN.value,
)
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={
"sub": "testuser@example.com",
"custom_roles": ["my-admin"],
},
user_info=existing_user,
prisma_client=prisma,
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
prisma.db.litellm_usertable.update.assert_not_called()
# ── VERIA-34 regression: PKCE state-to-session-cookie binding ───────────────
class TestPKCEStateCookieBinding:
"""The Generic SSO PKCE flow used the URL ``state`` parameter as a
cache-key for the PKCE ``code_verifier`` without binding the state to
the caller's browser. An attacker who pre-mints a state + cached
verifier could hand the link to a victim and capture the resulting
access token. Fix: set ``litellm_oauth_state`` HttpOnly cookie on
the redirect; verify the URL state matches the cookie before doing
the PKCE token exchange."""
@pytest.mark.asyncio
async def test_redirect_response_sets_oauth_state_cookie_when_pkce_enabled(self):
"""``get_generic_sso_redirect_response`` must set
``litellm_oauth_state`` on the redirect response when PKCE is on so
the callback can verify it later. The cookie must carry HttpOnly,
SameSite=Lax, and (because no http request was supplied to the
helper) the production-safe ``Secure`` default."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="https://idp.example.com/authorize?state=test-state-xyz"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "test-state-xyz",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="https://idp.example.com/authorize",
)
assert response is not None
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert (
cookie_str is not None
), f"litellm_oauth_state cookie not set; got: {cookie_headers}"
assert "test-state-xyz" in cookie_str
assert "HttpOnly" in cookie_str
assert "SameSite=lax" in cookie_str
# No incoming Request supplied → ``Secure`` defaults to True so a
# network observer on plain HTTP cannot read the state value.
assert "Secure" in cookie_str
@pytest.mark.asyncio
async def test_redirect_response_omits_oauth_state_cookie_when_pkce_disabled(
self,
):
"""Non-PKCE flows delegate to fastapi-sso's own session-cookie
binding; we do not set our cookie there because it would never be
validated (and could collide with a concurrent PKCE session in
the same browser)."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="https://idp.example.com/authorize?state=test-state-xyz"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "test-state-xyz",
"GENERIC_CLIENT_USE_PKCE": "false",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="https://idp.example.com/authorize",
)
assert response is not None
cookie_headers = response.headers.getlist("set-cookie")
assert not any(
"litellm_oauth_state=" in c for c in cookie_headers
), f"litellm_oauth_state cookie set on non-PKCE flow; got: {cookie_headers}"
@pytest.mark.asyncio
async def test_redirect_response_drops_secure_flag_for_http_dev(self):
"""When the incoming request is plain HTTP (local dev), ``Secure``
must be dropped so the browser will actually attach the cookie on
the callback hop."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="http://idp.local/authorize?state=local-dev-state"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
http_request = MagicMock(spec=Request)
http_request.url.scheme = "http"
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "local-dev-state",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="http://idp.local/authorize",
request=http_request,
)
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert cookie_str is not None
assert "Secure" not in cookie_str
@pytest.mark.asyncio
async def test_redirect_response_sets_secure_flag_behind_trusted_tls_terminating_proxy(
self, monkeypatch
):
"""Regression: litellm sees a plain-HTTP hop when TLS terminates at a reverse
proxy. The Secure flag must still be set when the direct peer is a configured
trusted proxy and it reports X-Forwarded-Proto: https -- but NOT from an
unconfigured/untrusted caller spoofing the same header (see the sibling test
below)."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="http://idp.internal/authorize?state=behind-proxy-state"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
proxied_request = MagicMock(spec=Request)
proxied_request.url.scheme = "http"
proxied_request.headers = {"X-Forwarded-Proto": "https"}
proxied_request.client = MagicMock()
proxied_request.client.host = "10.0.0.5"
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]},
)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "behind-proxy-state",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="http://idp.internal/authorize",
request=proxied_request,
)
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert cookie_str is not None
assert "Secure" in cookie_str
@pytest.mark.asyncio
async def test_redirect_response_ignores_spoofed_forwarded_proto_without_trust_config(
self, monkeypatch
):
"""The same X-Forwarded-Proto: https header must NOT flip Secure on when no
trusted-proxy config is present -- honoring it unconditionally would let any
client spoof the header and would not itself be the vulnerability the ticket
warns against."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="http://idp.internal/authorize?state=spoofed-state"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
spoofed_request = MagicMock(spec=Request)
spoofed_request.url.scheme = "http"
spoofed_request.headers = {"X-Forwarded-Proto": "https"}
spoofed_request.client = MagicMock()
spoofed_request.client.host = "203.0.113.5"
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "spoofed-state",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="http://idp.internal/authorize",
request=spoofed_request,
)
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert cookie_str is not None
assert "Secure" not in cookie_str
@pytest.mark.asyncio
async def test_pkce_callback_rejects_missing_cookie(self):
"""When PKCE is enabled and a code_verifier is in the cache, the
callback must reject a request that has no ``litellm_oauth_state``
cookie (browser-to-server binding missing)."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {
"state": "attacker-minted-state",
"code": "auth-code",
}
# No oauth_state cookie set → request.cookies.get returns None.
mock_request.cookies = {}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "attacker-cached-verifier",
"_pkce_cache_key": "pkce_verifier:attacker-minted-state",
}
),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
pytest.raises(ProxyException) as exc_info,
):
await get_generic_sso_response(
request=mock_request,
jwt_handler=MagicMock(spec=JWTHandler),
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert "state" in str(exc_info.value.message).lower()
@pytest.mark.asyncio
async def test_pkce_callback_rejects_state_cookie_mismatch(self):
"""The Login-CSRF shape: attacker mints state ``A``, victim's browser
carries cookie state ``B``. The callback must reject."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {
"state": "attacker-minted-state",
"code": "auth-code",
}
mock_request.cookies = {"litellm_oauth_state": "victim-browser-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:attacker-minted-state",
}
),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
pytest.raises(ProxyException) as exc_info,
):
await get_generic_sso_response(
request=mock_request,
jwt_handler=MagicMock(spec=JWTHandler),
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert "state" in str(exc_info.value.message).lower()
@pytest.mark.asyncio
async def test_pkce_callback_accepts_matching_state_cookie(self):
"""Happy path: URL state and cookie state match (the legitimate
flow where the same browser that started the redirect lands on
the callback) → the PKCE token exchange proceeds."""
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "matched-state", "code": "auth-code"}
mock_request.cookies = {"litellm_oauth_state": "matched-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:matched-state",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_pkce_token_exchange",
AsyncMock(
return_value={
"access_token": "tok",
"id_token": "id",
"sub": "user@example.com",
"email": "user@example.com",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_delete_pkce_verifier",
AsyncMock(),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
):
jwt_handler = MagicMock(spec=JWTHandler)
jwt_handler.get_team_ids_from_jwt.return_value = []
result, _, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=jwt_handler,
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
# State-cookie check passed, so the function got past the early
# ProxyException raise and produced an SSO result object.
assert result is not None
@pytest.mark.asyncio
@pytest.mark.parametrize("enable_sso_debug_value", [None, "false", "0"])
async def test_sso_debug_routes_return_404_unless_explicitly_enabled(enable_sso_debug_value):
"""
/sso/debug/login and /sso/debug/callback must 404 unless ENABLE_SSO_DEBUG is
explicitly set to a truthy value.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback, debug_sso_login
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
env = {"GENERIC_CLIENT_ID": "test_client_id"}
if enable_sso_debug_value is not None:
env["ENABLE_SSO_DEBUG"] = enable_sso_debug_value
with patch.dict(os.environ, env, clear=False):
if enable_sso_debug_value is None:
os.environ.pop("ENABLE_SSO_DEBUG", None)
with pytest.raises(HTTPException) as login_exc:
await debug_sso_login(mock_request)
with pytest.raises(HTTPException) as callback_exc:
await debug_sso_callback(mock_request)
assert login_exc.value.status_code == 404
assert callback_exc.value.status_code == 404
@pytest.mark.asyncio
async def test_debug_sso_callback_renders_full_jwt_claims():
"""
/sso/debug/callback should render the complete set of claims returned by the
IdP — both the raw userinfo response and the decoded access-token JWT — in
addition to the proxy-parsed OpenID fields. Bearer tokens must be stripped
even if a non-conforming IdP places them in its userinfo response.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
parsed_openid = CustomOpenID(
id="user_123",
email="philip@example.com",
first_name="Philip",
last_name="Schwartz",
display_name="Philip Schwartz",
provider="generic",
team_ids=["ord-engineering-high"],
user_role=None,
)
raw_userinfo_with_leaked_token = {
"sub": "user_123",
"email": "philip@example.com",
"team_id": "ord-engineering-high",
"team_alias": "ord-engineering-high",
"teams": ["ord-engineering-high"],
"roles": ["litellm.api.user"],
# Defense-in-depth: a non-conforming IdP could shove a bearer token
# into userinfo. The debug endpoint must strip it before rendering.
"access_token": "should-not-render",
"id_token": "should-not-render-either",
}
access_token_payload = {
"sub": "user_123",
"scope": "openid profile email",
"groups": ["litellm-users"],
}
async def fake_get_generic_sso_response(**kwargs):
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload, None
with (
patch.dict(
os.environ,
{"GENERIC_CLIENT_ID": "test_client_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response",
side_effect=fake_get_generic_sso_response,
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)),
):
# Microsoft / Google envs may leak in from other tests — ensure only
# the generic path runs.
for var in ("MICROSOFT_CLIENT_ID", "GOOGLE_CLIENT_ID"):
os.environ.pop(var, None)
response = await debug_sso_callback(mock_request)
body = response.body.decode()
# The embedded JSON payload drives the rendered page. Extract and parse it
# so we can assert on shape, not on cosmetic HTML details.
marker = "const ssoData = "
start = body.index(marker) + len(marker)
end = body.index(";", start)
while body[end - 1] not in "}]": # handle ';' inside string values
end = body.index(";", end + 1)
payload = json.loads(body[start:end])
assert set(payload.keys()) == {
"parsed_by_proxy",
"raw_claims",
"access_token_claims",
}
# Parsed OpenID fields are shown
assert payload["parsed_by_proxy"]["email"] == "philip@example.com"
assert payload["parsed_by_proxy"]["id"] == "user_123"
# Raw IdP claims surface fields the OpenID model drops (the original LIT-2838 ask)
assert payload["raw_claims"]["team_id"] == "ord-engineering-high"
assert payload["raw_claims"]["team_alias"] == "ord-engineering-high"
assert payload["raw_claims"]["teams"] == ["ord-engineering-high"]
assert payload["raw_claims"]["roles"] == ["litellm.api.user"]
# Defense-in-depth: bearer tokens must never appear in the rendered HTML
assert "access_token" not in payload["raw_claims"]
assert "id_token" not in payload["raw_claims"]
assert "should-not-render" not in body
# Decoded access-token JWT claims are surfaced
assert payload["access_token_claims"]["groups"] == ["litellm-users"]
@pytest.mark.asyncio
async def test_debug_sso_callback_handles_missing_raw_response():
"""
Microsoft and Google paths don't return a raw response or access-token
payload. The debug endpoint must still render successfully with empty
sections instead of crashing.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
parsed_openid = CustomOpenID(
id="user_456",
email="user@example.com",
first_name="Some",
last_name="User",
display_name="Some User",
provider="microsoft",
team_ids=[],
user_role=None,
)
async def fake_microsoft_callback(**kwargs):
return parsed_openid
with (
patch.dict(
os.environ,
{"MICROSOFT_CLIENT_ID": "test_microsoft_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch.object(
MicrosoftSSOHandler,
"get_microsoft_callback_response",
side_effect=fake_microsoft_callback,
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)),
):
for var in ("GENERIC_CLIENT_ID", "GOOGLE_CLIENT_ID"):
os.environ.pop(var, None)
response = await debug_sso_callback(mock_request)
assert response.status_code == 200
body = response.body.decode()
assert '"raw_claims": {}' in body
assert '"access_token_claims": {}' in body
assert "user@example.com" in body
# ── The debug page is where an operator lands when ID-JAG is failing ──────────
_GOOGLE_DEBUG_CLIENT_ID = "debug-google-client-id"
_GENERIC_DEBUG_CLIENT_ID = "debug-generic-client-id"
async def _render_debug_page(provider_env, id_jag_registered, force_inert=False):
"""Drive /sso/debug/callback and return the raw response body."""
from litellm.proxy.management_endpoints.ui_sso import GoogleSSOHandler, debug_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
parsed = {"sub": "user_123", "email": "u@example.com"}
async def fake_generic(**kwargs):
return parsed, {"sub": "user_123"}, {"scope": "openid"}, None
async def fake_google(**kwargs):
return parsed
stack = [
patch.dict(os.environ, {**provider_env, "ENABLE_SSO_DEBUG": "true"}, clear=False),
patch( # test-quality-ok: endpoint test stubs the upstream generic IdP boundary
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic
),
patch.object( # test-quality-ok: endpoint test stubs the upstream Google IdP boundary
GoogleSSOHandler, "get_google_callback_response", side_effect=fake_google
),
patch( # test-quality-ok: debug endpoint reads this module global without an injection seam
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=id_jag_registered),
),
patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: debug endpoint reads proxy globals
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: debug endpoint reads proxy DB
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: debug endpoint reads proxy globals
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)), # test-quality-ok: debug endpoint reads proxy globals
]
if force_inert:
stack.append(
patch( # test-quality-ok: force-inert reference isolates the endpoint's pre-change response
"litellm.proxy.management_endpoints.ui_sso.warn_if_id_jag_capture_gap",
AsyncMock(return_value=None),
)
)
with ExitStack() as es:
for ctx in stack:
es.enter_context(ctx)
for var in ("MICROSOFT_CLIENT_ID", "GOOGLE_CLIENT_ID", "GENERIC_CLIENT_ID", "SAML_IDP_METADATA_URL"):
if var not in provider_env:
os.environ.pop(var, None)
response = await debug_sso_callback(mock_request)
return response.body.decode()
@pytest.mark.asyncio
async def test_debug_page_logs_the_capture_gap_but_never_renders_it(caplog):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
body = await _render_debug_page({"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID}, id_jag_registered=True)
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert "google" in warnings[0]
assert "GENERIC_CLIENT_ID" in warnings[0]
assert "id_jag" not in body
assert "GENERIC_CLIENT_ID" not in body
@pytest.mark.asyncio
async def test_debug_page_is_byte_identical_when_the_provider_captures():
"""A deployment with no gap must get the page it got before this change, to the byte. The
comparison is against the endpoint with the diagnostic forced inert, not against a guess."""
with_feature = await _render_debug_page(
{"GENERIC_CLIENT_ID": _GENERIC_DEBUG_CLIENT_ID}, id_jag_registered=True
)
pre_change = await _render_debug_page(
{"GENERIC_CLIENT_ID": _GENERIC_DEBUG_CLIENT_ID},
id_jag_registered=True,
force_inert=True,
)
assert with_feature == pre_change
assert "id_jag" not in with_feature
@pytest.mark.asyncio
async def test_debug_page_is_byte_identical_when_no_id_jag_server_is_registered():
"""Most deployments run Google SSO and no id_jag server at all; their debug page must not
grow an ID-JAG section about a feature they do not use."""
with_feature = await _render_debug_page(
{"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID}, id_jag_registered=False
)
pre_change = await _render_debug_page(
{"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID},
id_jag_registered=False,
force_inert=True,
)
assert with_feature == pre_change
assert "id_jag" not in with_feature
@pytest.mark.asyncio
async def test_debug_page_survives_a_store_outage(monkeypatch, caplog):
"""The page's job is to render claims; an unreachable MCP table must cost it the annotation,
not the page."""
from litellm.proxy.management_endpoints.ui_sso import warn_if_id_jag_capture_gap
monkeypatch.setenv("GOOGLE_CLIENT_ID", _GOOGLE_DEBUG_CLIENT_ID)
retention_check = AsyncMock(side_effect=Exception("db down"))
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
assert await warn_if_id_jag_capture_gap(retention_enabled=retention_check) is None
retention_check.assert_awaited_once()
assert _id_jag_gap_warnings(caplog) == []
async def _render_legacy_login_page(env_overrides, general_settings):
from litellm.proxy.management_endpoints.ui_sso import google_login
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
with (
# snapshot os.environ so the mutations below are reverted on exit
patch.dict(os.environ, {}, clear=False),
patch("litellm.proxy.proxy_server.master_key", "sk-1234"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.proxy.proxy_server.general_settings", general_settings),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", None),
):
# No SSO provider configured, so /sso/key/generate renders the legacy
# username/password form rather than redirecting to an IdP.
for var in (
"MICROSOFT_CLIENT_ID",
"GOOGLE_CLIENT_ID",
"GENERIC_CLIENT_ID",
"LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT",
"UI_PASSWORD",
):
os.environ.pop(var, None)
os.environ.update(env_overrides)
return await google_login(request=mock_request)
@pytest.mark.asyncio
async def test_legacy_login_page_shows_credentials_hint_by_default():
"""Control: without the flag, the legacy page still discloses the hint."""
response = await _render_legacy_login_page(env_overrides={}, general_settings={})
body = response.body.decode()
assert response.status_code == 200
assert "Default Credentials" in body
assert "MASTER_KEY" in body
@pytest.mark.asyncio
async def test_legacy_login_page_hides_credentials_hint_via_env_flag():
"""
Regression: an anonymous GET /sso/key/generate must not disclose the
'admin / MASTER_KEY' default-credentials hint when
LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT is set. The legacy server-rendered
page previously ignored this flag while the new UI honored it.
"""
response = await _render_legacy_login_page(
env_overrides={"LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT": "true"},
general_settings={},
)
body = response.body.decode()
assert response.status_code == 200
assert "Default Credentials" not in body
assert "MASTER_KEY" not in body
# the login form itself must still render
assert 'name="username"' in body
@pytest.mark.asyncio
async def test_legacy_login_page_hides_credentials_hint_via_general_settings():
"""The flag is also honored from general_settings, matching the discovery endpoint."""
response = await _render_legacy_login_page(
env_overrides={},
general_settings={"hide_default_credentials_hint": True},
)
body = response.body.decode()
assert response.status_code == 200
assert "Default Credentials" not in body
assert "MASTER_KEY" not in body
@pytest.mark.asyncio
async def test_legacy_login_page_hides_credentials_hint_when_ui_password_set():
response = await _render_legacy_login_page(
env_overrides={"UI_PASSWORD": "s3cret-pass"},
general_settings={},
)
body = response.body.decode()
assert response.status_code == 200
assert "Default Credentials" not in body
assert "MASTER_KEY" not in body
assert 'name="username"' in body
@pytest.mark.asyncio
async def test_saml_callback_blocked_when_admin_ui_disabled():
"""An IdP-initiated assertion must not mint a UI session when the admin UI is
disabled; the ACS enforces DISABLE_ADMIN_UI like the SP-initiated login route."""
from litellm.proxy.management_endpoints.ui_sso import saml_callback
with patch.dict(os.environ, {"DISABLE_ADMIN_UI": "true"}):
response = await saml_callback(SimpleNamespace(cookies={}))
assert response.status_code == 200
assert "Admin UI is Disabled" in response.body.decode()
@pytest.mark.asyncio
async def test_saml_callback_enforces_free_sso_user_limit_after_validation():
"""An IdP-initiated assertion must not bypass the >5 free-SSO-user Enterprise gate
that /sso/key/generate enforces; the ACS re-checks it after validating the assertion,
so the entitlement DB query never runs on unvalidated input."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import saml_callback
call_order: list[str] = []
async def _fake_handle_acs(**kwargs):
call_order.append("validate")
return CustomOpenID(
id="dana@litellm.ai",
email="dana@litellm.ai",
first_name=None,
last_name=None,
display_name="dana",
picture=None,
provider="saml",
team_ids=[],
user_role=None,
)
async def _fake_count_billable_users():
call_order.append("count")
return 6
async def _stream():
yield b"SAMLResponse=signed-response"
request_double = SimpleNamespace(cookies={}, headers={}, stream=_stream)
with patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}), patch(
"litellm.proxy.proxy_server.premium_user", False
), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch(
"litellm.proxy.proxy_server.master_key", "sk-1234"
), patch(
"litellm.proxy.management_endpoints.sso.saml_sso.SAMLAuthHandler.handle_acs",
new=_fake_handle_acs,
), patch(
"litellm.repositories.user_repository.UserRepository.count_billable_users",
new=AsyncMock(side_effect=_fake_count_billable_users),
):
with pytest.raises(ProxyException) as exc:
await saml_callback(request_double)
assert str(exc.value.code) == "403"
assert call_order == ["validate", "count"]
@pytest.mark.asyncio
async def test_cli_poll_key_tolerates_missing_user_row():
"""The CLI poll must still mint the JWT when the user lookup raises,
e.g. the user row was created moments ago and a negative-cache window
from the pre-creation SSO existence check is still active on this pod."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_key = "cli-session-missing-user"
session_data = {
"user_id": "just-created-user",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
}
mock_cache = MagicMock(redis_cache=None)
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.missing.user"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
patch("litellm.proxy.proxy_server.prisma_client"),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
),
patch(
"litellm.proxy.auth.auth_checks.get_user_object",
new=AsyncMock(side_effect=ValueError("User doesn't exist in db. 'user_id'=just-created-user")),
),
):
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
assert result["key"] == mock_jwt_token
assert result["user_id"] == "just-created-user"
def _make_sso_callback_request(query_params: dict) -> MagicMock:
mock_request = MagicMock(spec=Request)
mock_request.query_params = query_params
return mock_request
@pytest.mark.asyncio
async def test_auth_callback_surfaces_oauth_error_with_description():
"""
Regression: when the IdP denies access it redirects back with
?error=...&error_description=... and no `code`. The callback must surface
that reason as a 401 instead of failing later on the missing `code` param.
"""
from litellm.proxy.management_endpoints.ui_sso import auth_callback
mock_request = _make_sso_callback_request(
{"error": "access_denied", "error_description": "User is not assigned to the client application"}
)
with pytest.raises(HTTPException) as exc_info:
await auth_callback(request=mock_request, state=None)
assert exc_info.value.status_code == 401
assert "access_denied" in str(exc_info.value.detail)
assert "User is not assigned to the client application" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_auth_callback_surfaces_oauth_error_without_description():
"""error_description is optional in the OAuth error response; the 401 detail must not render 'None'."""
from litellm.proxy.management_endpoints.ui_sso import auth_callback
mock_request = _make_sso_callback_request({"error": "access_denied"})
with pytest.raises(HTTPException) as exc_info:
await auth_callback(request=mock_request, state=None)
assert exc_info.value.status_code == 401
assert "access_denied" in str(exc_info.value.detail)
assert "None" not in str(exc_info.value.detail)
assert "error_description" not in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_auth_callback_without_oauth_error_proceeds_to_normal_flow():
"""Without an `error` query param the guard must not fire; the callback proceeds into the normal flow."""
from litellm.proxy.management_endpoints.ui_sso import auth_callback
mock_request = _make_sso_callback_request({"code": "some-auth-code"})
with patch("litellm.proxy.proxy_server.prisma_client", None):
with pytest.raises(HTTPException) as exc_info:
await auth_callback(request=mock_request, state=None)
assert exc_info.value.status_code == 500
assert "DB not connected" in str(exc_info.value.detail)
# ── SSO identity assertion capture + persist wiring (EMA) ─────────────────────
def _ema_id_token(sub: str = "u1") -> str:
import time as _time
import jwt as _pyjwt
return _pyjwt.encode(
{"iss": "https://idp.example.com", "sub": sub, "exp": int(_time.time()) + 3600},
"test-idp-signing-key-32-bytes-long-xxxx",
algorithm="HS256",
)
@pytest.mark.asyncio
async def test_pkce_arm_captures_sso_assertion():
"""The PKCE token exchange strips bearer fields from received_response for safety;
the typed assertion carrier must still capture id_token + refresh_token."""
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
id_token = _ema_id_token()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "matched-state", "code": "auth-code"}
mock_request.cookies = {"litellm_oauth_state": "matched-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:matched-state",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_pkce_token_exchange",
AsyncMock(
return_value={
"access_token": "tok",
"id_token": id_token,
"refresh_token": "rt_from_idp",
"sub": "user@example.com",
"email": "user@example.com",
}
),
),
patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
):
jwt_handler = MagicMock(spec=JWTHandler)
jwt_handler.get_team_ids_from_jwt.return_value = []
result, received_response, _, sso_assertion = await get_generic_sso_response(
request=mock_request,
jwt_handler=jwt_handler,
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert sso_assertion is not None
assert sso_assertion.id_token.get_secret_value() == id_token
assert sso_assertion.refresh_token is not None
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
# The sanitized received_response must still not carry bearer material.
assert "id_token" not in (received_response or {})
assert "refresh_token" not in (received_response or {})
@pytest.mark.asyncio
async def test_verify_and_process_arm_captures_sso_assertion():
"""The non-PKCE generic arm reads the raw bearer fields off the fastapi-sso client."""
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
id_token = _ema_id_token()
mock_request = MagicMock(spec=Request)
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
mock_sso_instance = MagicMock()
mock_sso_instance.verify_and_process = AsyncMock(
return_value={"sub": "u1", "email": "u@example.com"}
)
mock_sso_instance.access_token = None
mock_sso_instance.id_token = id_token
mock_sso_instance.refresh_token = "rt_from_idp"
mock_sso_class = MagicMock(return_value=mock_sso_instance)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "test_secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
},
):
with patch("fastapi_sso.sso.base.DiscoveryDocument"):
with patch(
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
_, _, _, sso_assertion = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id="test_client_id",
redirect_url="http://test.com/callback",
sso_jwt_handler=None,
)
assert sso_assertion is not None
assert sso_assertion.id_token.get_secret_value() == id_token
assert sso_assertion.refresh_token is not None
assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp"
@pytest.mark.asyncio
async def test_redirect_from_openid_persists_assertion_under_canonical_user_id():
"""The browser funnel persists the captured assertion AFTER canonical user
resolution, keyed by the user_id admission will later resolve (the key-generation
response user_id), not the raw IdP subject."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
)
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
retain_mock = AsyncMock()
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.proxy.proxy_server.user_custom_sso", None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
AsyncMock(
return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"}
),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id",
AsyncMock(return_value="internal_user"),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
retain_mock,
),
):
response = await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=CustomOpenID(
id="raw-idp-subject",
email="u@example.com",
first_name="U",
last_name="Ser",
display_name="U Ser",
provider="generic",
team_ids=[],
user_role=None,
),
request=mock_request,
received_response=None,
generic_client_id="cid",
ui_access_mode=None,
access_token_payload=None,
jwt_handler=None,
sso_assertion=assertion,
)
retain_mock.assert_awaited_once_with(
user_id="canonical-user-id", assertion=assertion
)
assert response is not None
@pytest.mark.asyncio
async def test_cli_completion_persists_assertion_under_db_user_id():
"""The CLI funnel persists the captured assertion under the DB-resolved user_id."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
)
from litellm.proxy.management_endpoints.ui_sso import (
_complete_cli_sso_callback_session,
)
assertion = assertion_from_sso_login(_ema_id_token(), None)
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
user_info = MagicMock()
user_info.user_id = "cli-user-id"
user_info.user_role = "internal_user"
user_info.models = []
user_info.teams = []
retain_mock = AsyncMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=user_info),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=[]),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
return_value={},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
retain_mock,
),
):
response = await _complete_cli_sso_callback_session(
request=mock_request,
key="cli-login-id",
flow={},
result={"sub": "raw-idp-subject"},
parsed_openid_result={
"user_id": "raw-idp-subject",
"user_email": "u@example.com",
"user_role": None,
},
user_defined_values=None,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
cli_sso_session_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
sso_assertion=assertion,
)
retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion)
assert response.status_code == 200
def _id_jag_gap_warnings(caplog) -> list[str]:
return [
record.getMessage()
for record in caplog.records
if record.levelno == logging.WARNING and "oauth2_id_jag" in record.getMessage()
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"provider_env, expected_fragment",
[
({"GOOGLE_CLIENT_ID": "cid"}, "google"),
({"MICROSOFT_CLIENT_ID": "cid", "MICROSOFT_TENANT": "t"}, "microsoft"),
({}, "no SSO provider is configured"),
],
)
async def test_uncaptured_assertion_warns_when_an_id_jag_server_is_registered(
monkeypatch, caplog, provider_env, expected_fragment
):
"""A provider with no capture path leaves ID-JAG permanently broken, and the only place
that is knowable is the login itself; without this line the operator sees nothing at all."""
from litellm.proxy.management_endpoints.ui_sso import (
warn_if_id_jag_assertion_uncaptured,
)
for name in ("GOOGLE_CLIENT_ID", "MICROSOFT_CLIENT_ID", "GENERIC_CLIENT_ID", "SAML_IDP_METADATA_URL"):
monkeypatch.delenv(name, raising=False)
for name, value in provider_env.items():
monkeypatch.setenv(name, value)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=True))
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert expected_fragment in str(warnings[0])
@pytest.mark.asyncio
async def test_generic_provider_that_returned_no_id_token_still_warns(monkeypatch, caplog):
"""Generic OIDC has a capture path, so there is no configuration gap to report; the login
still handed the id_jag arm nothing, and that must not pass silently."""
from litellm.proxy.management_endpoints.ui_sso import (
warn_if_id_jag_assertion_uncaptured,
)
for name in ("GOOGLE_CLIENT_ID", "MICROSOFT_CLIENT_ID", "SAML_IDP_METADATA_URL"):
monkeypatch.delenv(name, raising=False)
monkeypatch.setenv("GENERIC_CLIENT_ID", "cid")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=True))
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert "no usable id_token" in str(warnings[0])
@pytest.mark.asyncio
async def test_no_warning_when_the_assertion_was_captured(monkeypatch, caplog):
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
assertion_from_sso_login,
)
from litellm.proxy.management_endpoints.ui_sso import (
warn_if_id_jag_assertion_uncaptured,
)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
assertion = assertion_from_sso_login(_ema_id_token(), None)
assert assertion is not None
retention_mock = AsyncMock(return_value=True)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await warn_if_id_jag_assertion_uncaptured(assertion, retention_enabled=retention_mock)
assert _id_jag_gap_warnings(caplog) == []
retention_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_no_warning_when_no_id_jag_server_is_registered(monkeypatch, caplog):
"""Most deployments never register one; a warning about ID-JAG on every login there would
be pure noise and would train operators to ignore it."""
from litellm.proxy.management_endpoints.ui_sso import (
warn_if_id_jag_assertion_uncaptured,
)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=False))
assert _id_jag_gap_warnings(caplog) == []
@pytest.mark.asyncio
async def test_store_outage_does_not_break_the_login(monkeypatch, caplog):
from litellm.proxy.management_endpoints.ui_sso import (
warn_if_id_jag_assertion_uncaptured,
)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
assert (
await warn_if_id_jag_assertion_uncaptured(
None, retention_enabled=AsyncMock(side_effect=Exception("db down"))
)
is None
)
assert _id_jag_gap_warnings(caplog) == []
@pytest.mark.asyncio
async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
"""Wiring: the browser login path must reach the diagnostic, not just define it."""
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()
),
patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.premium_user", False), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.user_custom_sso", None), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.redis_usage_cache", None), # test-quality-ok: endpoint reads proxy globals
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: endpoint reads proxy globals
patch( # test-quality-ok: endpoint test stubs key generation at its module boundary
"litellm.proxy.proxy_server.generate_key_helper_fn",
AsyncMock(return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"}),
),
patch( # test-quality-ok: endpoint test stubs the user database lookup
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=None),
),
patch( # test-quality-ok: endpoint test stubs the admin database lookup
"litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id",
AsyncMock(return_value="internal_user"),
),
patch( # test-quality-ok: endpoint test stubs assertion persistence
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
patch( # test-quality-ok: endpoint reads this module global without an injection seam
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=True),
),
):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=CustomOpenID(
id="raw-idp-subject",
email="u@example.com",
first_name="U",
last_name="Ser",
display_name="U Ser",
provider="google",
team_ids=[],
user_role=None,
),
request=mock_request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
access_token_payload=None,
jwt_handler=None,
sso_assertion=None,
)
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert "google" in str(warnings[0])
@pytest.mark.asyncio
async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
"""Wiring: the CLI login path shares the gap, so it must share the diagnostic."""
from litellm.proxy.management_endpoints.ui_sso import (
_complete_cli_sso_callback_session,
)
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "cid")
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
user_info = MagicMock()
user_info.user_id = "cli-user-id"
user_info.user_role = "internal_user"
user_info.models = []
user_info.teams = []
with (
patch( # test-quality-ok: endpoint test stubs the user database lookup
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=user_info),
),
patch( # test-quality-ok: endpoint test stubs CLI team lookup
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=[]),
),
patch( # test-quality-ok: endpoint test stubs attribution metadata
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
return_value={},
),
patch( # test-quality-ok: endpoint test stubs assertion persistence
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
patch( # test-quality-ok: endpoint reads this module global without an injection seam
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=True),
),
):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await _complete_cli_sso_callback_session(
request=mock_request,
key="cli-login-id",
flow={},
result={"sub": "raw-idp-subject"},
parsed_openid_result={
"user_id": "raw-idp-subject",
"user_email": "u@example.com",
"user_role": None,
},
user_defined_values=None,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
cli_sso_session_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
sso_assertion=None,
)
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert "microsoft" in str(warnings[0])
def _cli_callback_kwargs(flow):
return {
"request": _cli_callback_request(),
"key": "cli-login-id",
"flow": flow,
"result": {"sub": "raw-idp-subject"},
"parsed_openid_result": {
"user_id": "raw-idp-subject",
"user_email": "u@example.com",
"user_role": None,
},
"user_defined_values": None,
"prisma_client": MagicMock(),
"user_api_key_cache": MagicMock(),
"cli_sso_session_cache": MagicMock(),
"proxy_logging_obj": MagicMock(),
}
def _cli_callback_request():
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
return mock_request
def _cli_callback_user_info(teams):
user_info = MagicMock()
user_info.user_id = "cli-user-id"
user_info.user_role = "internal_user"
user_info.models = ["personal-only"]
user_info.teams = teams
return user_info
@pytest.mark.asyncio
async def test_cli_completion_drops_teams_whose_rows_no_longer_exist():
"""A membership pointing at a deleted team must not be offered for selection.
Deleting an organization removes its team rows but leaves the user's membership
behind. If that dead team still reached the session, it would be auto-selected
for a single-team user, its grants could never resolve, and every future login
would be refused with no way for the user to recover.
"""
from litellm.proxy.management_endpoints.ui_sso import (
CliSsoTeamDetail,
_complete_cli_sso_callback_session,
)
live_detail = CliSsoTeamDetail(
team_id="team-live", team_alias="Live", team_models=("gpt-4.1",)
)
flow = {}
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=_cli_callback_user_info(["team-live", "team-deleted"])),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=(live_detail,)),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
return_value={},
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
):
response = await _complete_cli_sso_callback_session(**_cli_callback_kwargs(flow))
assert response.status_code == 200
assert flow["session_data"]["teams"] == ["team-live"]
assert [d["team_id"] for d in flow["session_data"]["team_details"]] == ["team-live"]
@pytest.mark.asyncio
async def test_cli_completion_fails_the_login_when_team_lookup_fails():
"""A lookup failure must fail the login instead of caching a teamless session.
Silently dropping every team here would hand a team-bound user a session with
their personal allowlist, which is the same "unknown grant treated as a real
grant" bug in a quieter form.
"""
from litellm.proxy.management_endpoints.ui_sso import (
_complete_cli_sso_callback_session,
)
flow = {}
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
AsyncMock(return_value=_cli_callback_user_info(["team-live"])),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
):
with pytest.raises(HTTPException) as exc_info:
await _complete_cli_sso_callback_session(**_cli_callback_kwargs(flow))
assert exc_info.value.status_code == 500
assert "session_data" not in flow
class TestSameOriginReturnPath:
"""The same-origin relative return_to arm added for the MCP gateway DCR authorize
round-trip: only strictly relative paths qualify, so login can never redirect the
browser off the gateway origin."""
def test_accepts_relative_paths(self):
from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path
assert _is_same_origin_return_path("/authorize?client_id=llm_dcrc_x&state=s") is True
assert _is_same_origin_return_path("/some_server/authorize") is True
def test_rejects_absolute_protocol_relative_and_backslash_paths(self):
from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path
assert _is_same_origin_return_path("https://evil.example.com/authorize") is False
assert _is_same_origin_return_path("//evil.example.com/authorize") is False
assert _is_same_origin_return_path("/\\evil.example.com") is False
assert _is_same_origin_return_path("javascript:alert(1)") is False
assert _is_same_origin_return_path("") is False
def _make_https_request() -> Request:
request = MagicMock(spec=Request)
request.url.scheme = "https"
request.headers = {}
request.client = MagicMock()
request.client.host = "203.0.113.5"
return request
def _make_http_request() -> Request:
request = MagicMock(spec=Request)
request.url.scheme = "http"
request.headers = {}
request.client = MagicMock()
request.client.host = "203.0.113.5"
return request
class TestPersistReturnToCookieSharedHelper:
"""The single shared return_to helper used by EVERY sign-in branch (SSO / Okta / generic AND the
username/password form). It must be best-effort and NEVER raise — a bad return_to can never block
sign-in. Regression: the password form previously 400'd because it called _validate_return_to
directly (which raises for a non-matching absolute return_to when control_plane_url is set)."""
@staticmethod
def _cookie(resp) -> str:
return resp.headers.get("set-cookie", "")
def test_sets_cookie_for_same_origin_relative_path(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
resp = Response()
_persist_return_to_cookie(resp, "/mcp/authorize?client_id=llm_dcrc_abc", _make_https_request())
assert "litellm_cp_return_to=" in self._cookie(resp)
def test_bad_absolute_with_control_plane_configured_does_not_raise_and_is_not_stored(self, monkeypatch):
"""THE regression: a non-matching absolute return_to with control_plane_url set must NOT raise
(it did, blocking the login form) and must NOT be stored — sign-in proceeds."""
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"}
)
resp = Response()
_persist_return_to_cookie(resp, "https://evil.example.com/steal", _make_https_request()) # must not raise
assert "litellm_cp_return_to=" not in self._cookie(resp)
def test_none_return_to_is_a_noop(self):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
resp = Response()
_persist_return_to_cookie(resp, None, _make_https_request())
assert "litellm_cp_return_to=" not in self._cookie(resp)
def test_control_plane_matching_absolute_is_stored(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"}
)
resp = Response()
_persist_return_to_cookie(resp, "https://cp.example.com/ui?page=models", _make_https_request())
assert "litellm_cp_return_to=" in self._cookie(resp)
def test_cookie_is_secure_and_httponly_over_https(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
resp = Response()
_persist_return_to_cookie(resp, "/mcp/authorize", _make_https_request())
cookie = self._cookie(resp)
assert "Secure" in cookie
assert "HttpOnly" in cookie
assert "SameSite=lax" in cookie
def test_cookie_is_not_secure_over_plain_http_direct(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
resp = Response()
_persist_return_to_cookie(resp, "/mcp/authorize", _make_http_request())
assert "Secure" not in self._cookie(resp)
def test_cookie_is_secure_behind_trusted_tls_terminating_proxy(self, monkeypatch):
"""Regression for the reported bug: TLS terminates at a reverse proxy, litellm only
sees a plain-HTTP hop, but a trusted X-Forwarded-Proto: https must still mark the
cookie Secure."""
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]},
)
resp = Response()
request = _make_http_request()
request.client.host = "10.0.0.5"
request.headers = {"X-Forwarded-Proto": "https"}
_persist_return_to_cookie(resp, "/mcp/authorize", request)
assert "Secure" in self._cookie(resp)
class TestSessionTokenCookie:
"""Regression tests for the ``token`` session cookie set by every sign-in path
(username/password login, SSO callback, the CLI /v2, /v3 login exchange helpers).
It was previously set with no Secure/HttpOnly/SameSite attributes at all -- always
sent over plain HTTP and readable by any script on the page. HttpOnly must stay off
deliberately: the dashboard reads this cookie via document.cookie."""
@staticmethod
def _cookie(resp) -> str:
return resp.headers.get("set-cookie", "")
def test_secure_over_direct_https(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
resp = Response()
set_session_token_cookie(resp, _make_https_request(), "jwt-token-value")
cookie = self._cookie(resp)
assert "token=jwt-token-value" in cookie
assert "Secure" in cookie
assert "SameSite=lax" in cookie
assert "HttpOnly" not in cookie
def test_not_secure_over_direct_http(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
resp = Response()
set_session_token_cookie(resp, _make_http_request(), "jwt-token-value")
assert "Secure" not in self._cookie(resp)
def test_secure_behind_trusted_tls_terminating_proxy(self, monkeypatch):
"""THE regression: TLS terminates at a reverse proxy, litellm only sees a
plain-HTTP hop, but the session cookie must still be marked Secure when the
operator has configured a trusted proxy that reports X-Forwarded-Proto: https."""
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]},
)
request = _make_http_request()
request.client.host = "10.0.0.5"
request.headers = {"X-Forwarded-Proto": "https"}
resp = Response()
set_session_token_cookie(resp, request, "jwt-token-value")
assert "Secure" in self._cookie(resp)
def test_untrusted_spoofed_forwarded_proto_is_ignored(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
request = _make_http_request()
request.headers = {"X-Forwarded-Proto": "https"}
resp = Response()
set_session_token_cookie(resp, request, "jwt-token-value")
assert "Secure" not in self._cookie(resp)
def test_proxy_base_url_https_overrides_literal_http_scheme(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import set_session_token_cookie
monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com")
resp = Response()
set_session_token_cookie(resp, _make_http_request(), "jwt-token-value")
assert "Secure" in self._cookie(resp)
@pytest.mark.asyncio
@pytest.mark.parametrize("trusted", [False, True])
@pytest.mark.parametrize("storage_available", [False, True])
async def test_cli_sign_in_enrolls_only_verified_subjects_before_completing(
monkeypatch: pytest.MonkeyPatch, trusted: bool, storage_available: bool
) -> None:
from typing import Final
from litellm.proxy.management_endpoints import ui_sso
from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
flow: Final[dict[str, object]] = {}
kwargs: Final = _cli_callback_kwargs(flow)
subject: Final = MicrosoftInteractiveSubject(issuer="issuer", tenant_id="tenant", oid="subject")
kwargs["request"].scope = {"litellm_microsoft_interactive_subject": subject if trusted else subject.model_dump()}
table: Final = kwargs["prisma_client"].writer_db.litellm_verifiedsubject
table.upsert = AsyncMock(
return_value=SimpleNamespace(kind="human", user_id="cli-user-id", verified_via="sso_interactive"),
side_effect=None if storage_available else RuntimeError("storage unavailable"),
)
monkeypatch.setattr(ui_sso, "get_user_info_from_db", AsyncMock(return_value=_cli_callback_user_info([])))
monkeypatch.setattr(ui_sso, "fetch_cli_sso_team_details", AsyncMock(return_value=()))
monkeypatch.setattr(ui_sso, "retain_sso_identity_assertion_for_ema", AsyncMock())
if trusted and not storage_available:
with pytest.raises(HTTPException) as error:
await ui_sso._complete_cli_sso_callback_session(**kwargs)
assert error.value.status_code == 503
assert "sso_complete" not in flow
return
response: Final = await ui_sso._complete_cli_sso_callback_session(**kwargs)
assert response.status_code == 200
assert flow["session_data"]["user_id"] == "cli-user-id"
if trusted:
table.upsert.assert_awaited_once_with(
where={"issuer_tenant_id_oid": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject"}},
data={"create": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject",
"user_id": "cli-user-id", "verified_via": "sso_interactive"}, "update": {}},
)
else:
table.upsert.assert_not_awaited()