feat(mcp): warn when an oauth2_id_jag server outruns the SSO provider's assertion capture (#35394)

* feat(mcp): warn when an oauth2_id_jag server outruns the SSO provider's assertion capture

Only the generic OIDC login path captures the IdP id_token that an oauth2_id_jag MCP
server spends as its RFC 8693 subject token. Under Google, Microsoft, SAML or no SSO at
all, registration succeeds and then every ID-JAG credential resolution fails for every
user, with nothing in the logs, the config or the API response to say why.

Report the mismatch from the two places it is knowable: when an oauth2_id_jag server is
created or updated through the management endpoint, and at SSO callback time when a login
hands the arm nothing while such a server is registered. Provider selection mirrors the
callback's precedence, so a generic client id sitting behind GOOGLE_CLIENT_ID does not
clear the warning.

* test(sso): update merged CLI diagnostic patch target

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

* feat(mcp): warn about the ID-JAG capture gap for config-declared servers and on the SSO debug page (#39350)

* feat(sso): surface the ID-JAG capture gap on the SSO debug page

/sso/debug/callback is where an operator lands when they are already trying to work out
why ID-JAG is failing, so the reason belongs on it. The annotation appears only when the
active SSO provider captures no identity assertion AND an oauth2_id_jag server is
registered for that gap to break; a deployment without both renders the page it rendered
before, byte for byte. Only the provider name and the remedy are rendered, never a
configured value, and an unreachable MCP table costs the page its annotation rather than
the page itself.

The payload carries the one mutable-ok in this work. Conditionally including a member of a
JSON document has to construct a mapping, and the rejected alternatives are recorded on the
helper so the next reader does not rediscover them.

Held out of the diagnosability PR deliberately: that PR is already reviewed and green, and
this surface ships with the remaining config-load warning as one follow-up.

* feat(mcp): warn at config load when an oauth2_id_jag server outruns the SSO provider's assertion capture

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

* refactor(sso): trim comments on the ID-JAG debug page diagnostic

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

* fix(sso): clean up merged imports

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

* fix(sso): satisfy type discipline for diagnostic payload

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

* style(sso): keep the optional ID-JAG payload member on one line for ruff format

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

* fix(sso): use Python 3.10-compatible assert_never

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

---------

Co-authored-by: Yassin Kortam <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(sso): keep the ID-JAG capture-gap diagnostic out of the unauthenticated debug page

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

* test(sso): inject the retention check and log via caplog so the ID-JAG tests pass the test-quality gate

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

* test(sso): keep the debug-page outage test on the capture-gap path

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

* fix(sso): annotate the retention check type alias

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

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yassin Kortam 2026-09-05 12:43:09 -07:00 • committed by GitHub
parent 110f654f34
commit 17e13126cc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 899 additions and 9 deletions

View file

@ -159,6 +159,9 @@ from litellm.proxy._types import (
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
id_jag_assertion_capture_gap_at_startup,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging, get_server_root_path
from litellm.repositories.table_repositories import MCPServerRepository
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -1382,6 +1385,20 @@ def _warn_internal_delegate_pkce_if_applicable(server: MCPServer, *, source: str
)
def _warn_config_id_jag_server_outruns_sso(server: MCPServer) -> None:
if server.auth_type != MCPAuth.oauth2_id_jag:
return
gap: Final = id_jag_assertion_capture_gap_at_startup()
if gap is None:
return
verbose_logger.warning(
"MCP server %r (id=%s, source=config) is declared with auth_type=oauth2_id_jag, but %s.",
get_server_prefix(server),
server.server_id,
gap,
)
def _deserialize_json_dict(data: str | _StringMap | None) -> dict[str, str] | None:
"""
Deserialize optional JSON mappings stored in the database.
@ -2393,6 +2410,7 @@ class MCPServerManager:
)
self._assign_unique_short_prefix(new_server)
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,

View file

@ -64,6 +64,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
id_jag_assertion_capture_gap,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
@ -272,6 +275,22 @@ if MCP_AVAILABLE:
_validate_mcp_server_name_fields(payload)
_validate_upstream_token_header(payload)
def warn_if_id_jag_server_outruns_sso(server_id: str | None, auth_type: MCPAuth | str | None) -> None:
"""Registering an ``oauth2_id_jag`` server under an SSO provider that captures no IdP
identity assertion is a dead configuration: nothing here fails, and then every ID-JAG call
fails for every user with a message that only ever tells them to sign in again. Say it once,
at the moment the admin can still act on it."""
if auth_type != MCPAuth.oauth2_id_jag:
return
gap = id_jag_assertion_capture_gap()
if gap is None:
return
verbose_proxy_logger.warning(
"MCP server %s is registered with auth_type=oauth2_id_jag, but %s.",
server_id,
gap,
)
def stamp_omitted_oauth2_flow(payload: NewMCPServerRequest) -> None:
"""Fallback only: fill in oauth2_flow when an oauth2 create omits it.
@ -1623,6 +1642,8 @@ if MCP_AVAILABLE:
detail={"error": f"Error creating mcp server: {e}"},
)
warn_if_id_jag_server_outruns_sso(new_mcp_server.server_id, new_mcp_server.auth_type)
# Registry refresh is best-effort: the row is already committed, so a
# failure here (e.g. an unrelated malformed row in the table) must not
# surface as a 500 and orphan the created server, which would push the
@ -2726,6 +2747,7 @@ if MCP_AVAILABLE:
status_code=status.HTTP_404_NOT_FOUND,
detail={"error": f"MCP Server not found, passed server_id={payload.server_id}"},
)
warn_if_id_jag_server_outruns_sso(mcp_server_record_updated.server_id, mcp_server_record_updated.auth_type)
await global_mcp_server_manager.update_server(mcp_server_record_updated)
# Ensure registry is up to date by reloading from database

View file

@ -0,0 +1,81 @@
"""Whether the SSO provider the login callback dispatches to can capture an IdP identity assertion.
An ``oauth2_id_jag`` MCP server spends the ``id_token`` captured at SSO login as its RFC 8693
subject token. Only the generic OIDC login path reaches a token response the gateway retains one
from, so a deployment whose SSO runs through Google, Microsoft or SAML never stores an assertion
and every store-sourced ID-JAG exchange fails for every user, however many times they sign in.
Neither side can see that alone: the MCP registration knows nothing about SSO and the login knows
nothing about MCP. This module is the one shared answer both warn from.
"""
from __future__ import annotations
import os
from enum import Enum
from typing_extensions import assert_never
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
_GENERIC_OIDC_REMEDY = (
"Point SSO at the generic OIDC provider (GENERIC_CLIENT_ID), the one login path whose token "
"response the gateway retains an id_token from"
)
class ActiveSSOProvider(str, Enum):
google = "google"
microsoft = "microsoft"
generic = "generic"
saml = "saml"
none = "none"
def active_sso_provider() -> ActiveSSOProvider:
"""The provider the SSO callback will dispatch to.
Mirrors the callback's precedence rather than reporting everything configured: an environment
carrying both GOOGLE_CLIENT_ID and GENERIC_CLIENT_ID runs the Google branch, so it must report
Google. Presence is judged the way the callback judges it, so a client id set to the empty
string still selects that branch here.
"""
if os.getenv("GOOGLE_CLIENT_ID") is not None:
return ActiveSSOProvider.google
if os.getenv("MICROSOFT_CLIENT_ID") is not None:
return ActiveSSOProvider.microsoft
if os.getenv("GENERIC_CLIENT_ID") is not None:
return ActiveSSOProvider.generic
if SAMLAuthHandler.is_saml_configured():
return ActiveSSOProvider.saml
return ActiveSSOProvider.none
def id_jag_assertion_capture_gap() -> str | None:
"""Why ID-JAG cannot work under the active SSO provider, phrased for an operator reading a log,
or ``None`` when that provider does capture an assertion."""
provider = active_sso_provider()
match provider:
case ActiveSSOProvider.generic:
return None
case ActiveSSOProvider.none:
return (
"no SSO provider is configured, so no IdP identity assertion is ever captured and "
f"ID-JAG credential resolution fails for every user. {_GENERIC_OIDC_REMEDY}"
)
case ActiveSSOProvider.google | ActiveSSOProvider.microsoft | ActiveSSOProvider.saml:
return (
f"the active SSO provider ({provider.value}) has no identity-assertion capture path, so no "
"IdP id_token is ever stored and ID-JAG credential resolution fails for every user no matter "
f"how often they sign in. {_GENERIC_OIDC_REMEDY}"
)
case _:
assert_never(provider)
def id_jag_assertion_capture_gap_at_startup() -> str | None:
"""Config load runs before SSO settings stored in the database are reconciled into the process
environment, so an unresolved provider at that point is not yet a gap; the SSO callback reports it
once a login happens."""
if active_sso_provider() is ActiveSSOProvider.none:
return None
return id_jag_assertion_capture_gap()

View file

@ -16,7 +16,7 @@ import json
import os
import re
import secrets
from collections.abc import Mapping, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from copy import deepcopy
from html import escape
from types import MappingProxyType
@ -29,6 +29,7 @@ from typing import (
NoReturn,
Optional,
Protocol,
TypeAlias,
Union,
cast,
overload,
@ -70,6 +71,7 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
SSOIdentityAssertion,
assertion_from_sso_login,
ema_assertion_retention_enabled,
retain_sso_identity_assertion_for_ema,
)
from litellm.proxy._types import (
@ -105,6 +107,9 @@ from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
id_jag_assertion_capture_gap,
)
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
from litellm.proxy.management_endpoints.sso_helper_utils import (
check_is_admin_only_access,
@ -1677,6 +1682,46 @@ async def get_generic_sso_response(
return result or {}, received_response, access_token_payload, sso_assertion
RetentionCheck: TypeAlias = Callable[[], Awaitable[bool]] # mutable-ok: Callable parameter syntax
async def warn_if_id_jag_assertion_uncaptured(
assertion: SSOIdentityAssertion | None, *, retention_enabled: RetentionCheck | None = None
) -> None:
"""Say, at the one moment it is knowable, that this login gave an ``oauth2_id_jag`` server
nothing to spend. Without it the operator only ever sees the per-request failure, which cannot
tell a user who has never signed in from a provider that will never capture. Kept strictly
diagnostic: a store outage is swallowed, since a login must not fail over a log line."""
if assertion is not None:
return
try:
check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
if not await check():
return
except Exception as exc: # noqa: BLE001 # diagnostics must never break the login
verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers after SSO login: %s", exc)
return
gap: Final = id_jag_assertion_capture_gap()
verbose_proxy_logger.warning(
"SSO login captured no IdP identity assertion while an oauth2_id_jag MCP server is registered: %s",
gap if gap is not None else "the identity provider's token response carried no usable id_token",
)
async def warn_if_id_jag_capture_gap(*, retention_enabled: RetentionCheck | None = None) -> None:
gap: Final = id_jag_assertion_capture_gap()
if gap is None:
return
try:
check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
if not await check():
return
except Exception as exc: # noqa: BLE001 # diagnostics must never break the page they annotate
verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers: %s", exc)
return
verbose_proxy_logger.warning("SSO debug callback ran with an oauth2_id_jag capture gap: %s", gap)
async def create_team_member_add_task(team_id, user_info):
"""Create a task for adding a member to a team."""
try:
@ -2269,6 +2314,7 @@ async def _complete_cli_sso_callback_session(
raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
teams: list[str] = []
if hasattr(user_info, "teams") and user_info.teams:
@ -3599,6 +3645,7 @@ class SSOAuthenticationHandler:
if isinstance(user_id, str) and user_id:
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
disabled_non_admin_personal_key_creation: Final = get_disabled_non_admin_personal_key_creation()
litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
@ -4733,6 +4780,7 @@ async def debug_sso_callback(request: Request):
safe_raw_claims: Final = {k: v for k, v in (received_response or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
safe_access_token_claims = {k: v for k, v in (access_token_payload or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
await warn_if_id_jag_capture_gap()
sso_payload: Final = {
"parsed_by_proxy": filtered_result,
"raw_claims": safe_raw_claims,

View file

@ -461,6 +461,30 @@ class TestMCPServerManager:
base.update(overrides)
return {"m2mserver": base}
def _id_jag_config(self):
return {
"idjag_server": {
"url": "https://example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2_id_jag,
"client_id": "cid",
"client_secret": "csec",
"token_exchange_endpoint": "https://idp.example.com/token",
"id_jag_resource_token_endpoint": "https://resource.example.com/token",
"id_jag_resource": "https://resource.example.com",
}
}
def _clear_sso_env(self, monkeypatch):
for env_var in (
"GOOGLE_CLIENT_ID",
"MICROSOFT_CLIENT_ID",
"GENERIC_CLIENT_ID",
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
):
monkeypatch.delenv(env_var, raising=False)
@pytest.mark.parametrize("value", ["1", "true", "TRUE", "yes", "on"])
def test_mcp_oauth_discovery_on_startup_true_values(self, value):
with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": value}):
@ -1130,6 +1154,72 @@ class TestMCPServerManager:
server = next(iter(manager.config_mcp_servers.values()))
assert server.oauth2_flow is None
@pytest.mark.asyncio
async def test_load_servers_from_config_warns_for_id_jag_with_google_sso(self, monkeypatch, caplog):
self._clear_sso_env(monkeypatch)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
warnings = [message for message in caplog.messages if "oauth2_id_jag" in message]
assert len(warnings) == 1
assert "idjag_server" in warnings[0]
assert "GENERIC_CLIENT_ID" in warnings[0]
@pytest.mark.asyncio
async def test_load_servers_from_config_does_not_warn_for_id_jag_without_sso(self, monkeypatch, caplog):
self._clear_sso_env(monkeypatch)
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
assert not any("oauth2_id_jag" in message for message in caplog.messages)
@pytest.mark.asyncio
async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, monkeypatch, caplog):
self._clear_sso_env(monkeypatch)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
manager = MCPServerManager()
config = {
"api_key_server": {
"url": "https://example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.api_key,
"auth_value": "upstream-secret",
}
}
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(config)
assert not any("oauth2_id_jag" in message for message in caplog.messages)
@pytest.mark.asyncio
async def test_load_servers_from_config_does_not_warn_for_id_jag_with_generic_sso(self, monkeypatch, caplog):
self._clear_sso_env(monkeypatch)
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
assert not any("oauth2_id_jag" in message for message in caplog.messages)
def _client_forwarded_config(self, auth_type, **overrides):
base = {
"url": "https://example.com/mcp",

View file

@ -0,0 +1,117 @@
import pytest
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
ActiveSSOProvider,
active_sso_provider,
id_jag_assertion_capture_gap,
id_jag_assertion_capture_gap_at_startup,
)
_SSO_ENV_VARS = (
"GOOGLE_CLIENT_ID",
"MICROSOFT_CLIENT_ID",
"GENERIC_CLIENT_ID",
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
)
@pytest.fixture(autouse=True)
def _isolated_sso_env(monkeypatch):
"""Every SSO selector is read from the process environment, so a value left behind by
another test would silently decide this one's answer."""
for name in _SSO_ENV_VARS:
monkeypatch.delenv(name, raising=False)
class TestActiveSSOProviderMirrorsTheCallback:
"""The gap warning is only as good as its agreement with the branch the login callback
actually takes, so provider selection is asserted branch by branch, including the
precedence that makes a co-configured generic client unreachable."""
def test_google_client_id_selects_google(self, monkeypatch):
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
assert active_sso_provider() is ActiveSSOProvider.google
def test_microsoft_client_id_selects_microsoft(self, monkeypatch):
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-cid")
assert active_sso_provider() is ActiveSSOProvider.microsoft
def test_generic_client_id_selects_generic(self, monkeypatch):
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
assert active_sso_provider() is ActiveSSOProvider.generic
def test_saml_metadata_selects_saml(self, monkeypatch):
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata")
assert active_sso_provider() is ActiveSSOProvider.saml
def test_nothing_configured_selects_none(self):
assert active_sso_provider() is ActiveSSOProvider.none
def test_google_outranks_a_co_configured_generic_client(self, monkeypatch):
"""The callback tests GOOGLE_CLIENT_ID first, so the generic arm never runs here and
no assertion is captured; reporting generic would clear a gap that is still open."""
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
assert active_sso_provider() is ActiveSSOProvider.google
def test_microsoft_outranks_a_co_configured_generic_client(self, monkeypatch):
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-cid")
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
assert active_sso_provider() is ActiveSSOProvider.microsoft
def test_generic_outranks_saml(self, monkeypatch):
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata")
assert active_sso_provider() is ActiveSSOProvider.generic
class TestIdJagAssertionCaptureGap:
def test_generic_oidc_has_no_gap(self, monkeypatch):
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
assert id_jag_assertion_capture_gap() is None
@pytest.mark.parametrize(
"env_var, provider_label",
[
("GOOGLE_CLIENT_ID", "google"),
("MICROSOFT_CLIENT_ID", "microsoft"),
("SAML_IDP_METADATA_URL", "saml"),
],
)
def test_non_capturing_provider_is_named_with_the_remedy(self, monkeypatch, env_var, provider_label):
monkeypatch.setenv(env_var, "configured")
gap = id_jag_assertion_capture_gap()
assert gap is not None
assert provider_label in gap
assert "GENERIC_CLIENT_ID" in gap
def test_no_sso_configured_reports_a_gap(self):
gap = id_jag_assertion_capture_gap()
assert gap is not None
assert "no SSO provider is configured" in gap
def test_google_beside_generic_still_reports_a_gap(self, monkeypatch):
"""The precedence trap in operator terms: adding a generic client id without removing
GOOGLE_CLIENT_ID does not fix the deployment, so the gap must not clear."""
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
gap = id_jag_assertion_capture_gap()
assert gap is not None
assert "google" in gap
class TestIdJagAssertionCaptureGapAtStartup:
def test_no_provider_at_startup_is_not_yet_a_gap(self):
assert id_jag_assertion_capture_gap_at_startup() is None
def test_google_provider_at_startup_reports_the_capture_gap(self, monkeypatch):
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
startup_gap = id_jag_assertion_capture_gap_at_startup()
callback_gap = id_jag_assertion_capture_gap()
assert startup_gap is not None
assert startup_gap == callback_gap
def test_generic_provider_at_startup_has_no_gap(self, monkeypatch):
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
assert id_jag_assertion_capture_gap_at_startup() is None

View file

@ -2,6 +2,7 @@ import os
import sys
import types
import json
import logging
from contextlib import ExitStack
from datetime import datetime, timedelta
from types import SimpleNamespace
@ -3840,6 +3841,155 @@ class TestAddMCPServerAtomicity:
mock_manager.reload_servers_from_database.assert_not_awaited()
class TestIdJagRegistrationWarnsAboutTheSSOGap:
"""An `oauth2_id_jag` server only ever works when the login path captures an IdP identity
assertion, and only the generic OIDC arm does. Registering one under Google or Microsoft
succeeds and then fails for every user on every call, so the mismatch has to be said at
registration time, while the admin is still looking at the configuration."""
@staticmethod
def _clear_sso_env(monkeypatch):
for name in (
"GOOGLE_CLIENT_ID",
"MICROSOFT_CLIENT_ID",
"GENERIC_CLIENT_ID",
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
):
monkeypatch.delenv(name, raising=False)
@staticmethod
def _id_jag_warnings(caplog) -> list[str]:
return [
record.getMessage()
for record in caplog.records
if record.levelno == logging.WARNING and "oauth2_id_jag" in record.getMessage()
]
@staticmethod
def _server_record(auth_type) -> LiteLLM_MCPServerTable:
record = generate_mock_mcp_server_db_record(server_id="ema-1", alias="ema")
record.auth_type = auth_type
return record
async def _run_create(self, monkeypatch, provider_env, auth_type, caplog):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
self._clear_sso_env(monkeypatch)
for name, value in provider_env.items():
monkeypatch.setenv(name, value)
mock_manager = MagicMock()
mock_manager.add_server = AsyncMock()
mock_manager.reload_servers_from_database = AsyncMock()
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch( # test-quality-ok: endpoint test stubs MCP server creation
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
AsyncMock(return_value=self._server_record(auth_type)),
),
patch( # test-quality-ok: endpoint reads the global MCP manager
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
),
):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await add_mcp_server(
payload=NewMCPServerRequest(
alias="ema",
url="https://ema.example.com/mcp",
transport=MCPTransport.http,
),
user_api_key_dict=generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
),
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"provider_env, expected_fragment",
[
({"GOOGLE_CLIENT_ID": "cid"}, "google"),
({"MICROSOFT_CLIENT_ID": "cid"}, "microsoft"),
({"SAML_IDP_METADATA_URL": "https://idp.example.com/metadata"}, "saml"),
({}, "no SSO provider is configured"),
],
)
async def test_create_warns_under_a_provider_that_captures_nothing(
self, monkeypatch, caplog, provider_env, expected_fragment
):
await self._run_create(monkeypatch, provider_env, MCPAuth.oauth2_id_jag, caplog)
warnings = self._id_jag_warnings(caplog)
assert len(warnings) == 1
assert expected_fragment in str(warnings[0])
assert "ema-1" in str(warnings[0])
@pytest.mark.asyncio
async def test_create_is_silent_under_generic_oidc(self, monkeypatch, caplog):
await self._run_create(monkeypatch, {"GENERIC_CLIENT_ID": "cid"}, MCPAuth.oauth2_id_jag, caplog)
assert self._id_jag_warnings(caplog) == []
@pytest.mark.asyncio
async def test_create_is_silent_for_other_auth_types(self, monkeypatch, caplog):
"""Nothing but the id_jag arm sources credentials from a stored SSO assertion, so no
other server registered under Google has anything to warn about."""
await self._run_create(monkeypatch, {"GOOGLE_CLIENT_ID": "cid"}, MCPAuth.api_key, caplog)
assert self._id_jag_warnings(caplog) == []
@pytest.mark.asyncio
async def test_update_to_id_jag_warns(self, monkeypatch, caplog):
"""Switching an existing server onto id_jag opens the same gap a create does."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
edit_mcp_server,
)
self._clear_sso_env(monkeypatch)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
mock_manager = MagicMock()
mock_manager.update_server = AsyncMock()
mock_manager.reload_servers_from_database = AsyncMock()
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch( # test-quality-ok: endpoint test stubs the MCP server lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=self._server_record(MCPAuth.api_key)),
),
patch( # test-quality-ok: endpoint test stubs MCP server updates
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
AsyncMock(return_value=self._server_record(MCPAuth.oauth2_id_jag)),
),
patch( # test-quality-ok: endpoint test stubs credential cleanup
"litellm.proxy.management_endpoints.mcp_management_endpoints.purge_user_oauth_credentials_for_server",
AsyncMock(return_value=0),
),
patch( # test-quality-ok: endpoint reads the global MCP manager
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
),
):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await edit_mcp_server(
payload=UpdateMCPServerRequest(server_id="ema-1", auth_type=MCPAuth.oauth2_id_jag),
user_api_key_dict=generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
),
)
warnings = self._id_jag_warnings(caplog)
assert len(warnings) == 1
assert "google" in str(warnings[0])
class TestHealthCheckServers:
"""Test suite for health check servers endpoint"""

View file

@ -1,17 +1,16 @@
import asyncio
import json
import logging
import os
from contextlib import asynccontextmanager
from contextlib import ExitStack, asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException, Request
from litellm._uuid import uuid
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
@ -1615,8 +1614,8 @@ async def test_get_generic_sso_response_with_empty_headers():
async def test_get_generic_sso_response_includes_token_claims_when_enabled(monkeypatch):
import jwt as pyjwt
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
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)
@ -2321,10 +2320,10 @@ class TestCustomUISSO:
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
@ -2400,6 +2399,7 @@ class TestCustomUISSO:
from litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
)
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
mock_request = MagicMock(spec=Request)
@ -2436,10 +2436,10 @@ class TestCustomUISSO:
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
@ -8167,6 +8167,128 @@ async def test_debug_sso_callback_handles_missing_raw_response():
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, 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
@ -8261,8 +8383,8 @@ async def test_saml_callback_enforces_free_sso_user_limit_after_validation():
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.ui_sso import saml_callback
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import saml_callback
call_order: list[str] = []
@ -8681,6 +8803,248 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
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.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.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(),