mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
* 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>
9661 lines
369 KiB
Python
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"><script>" 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()
|