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>
This commit is contained in:
yassin 2026-09-04 01:55:04 +00:00
parent 9c90998cb1
commit 237adf32e5
3 changed files with 183 additions and 204 deletions

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
@ -1681,7 +1681,12 @@ async def get_generic_sso_response(
return result or {}, received_response, access_token_payload, sso_assertion
async def warn_if_id_jag_assertion_uncaptured(assertion: SSOIdentityAssertion | None) -> None:
RetentionCheck = Callable[[], Awaitable[bool]]
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
@ -1689,7 +1694,8 @@ async def warn_if_id_jag_assertion_uncaptured(assertion: SSOIdentityAssertion |
if assertion is not None:
return
try:
if not await ema_assertion_retention_enabled():
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)
@ -1701,17 +1707,18 @@ async def warn_if_id_jag_assertion_uncaptured(assertion: SSOIdentityAssertion |
)
async def warn_if_id_jag_capture_gap() -> None:
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:
if not await ema_assertion_retention_enabled():
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 ID-JAG capture gap: %s", gap)
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):

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
@ -3858,8 +3859,12 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap:
monkeypatch.delenv(name, raising=False)
@staticmethod
def _id_jag_warnings(logger_mock) -> list:
return [call for call in logger_mock.warning.call_args_list if "oauth2_id_jag" in str(call)]
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:
@ -3867,7 +3872,7 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap:
record.auth_type = auth_type
return record
async def _run_create(self, monkeypatch, provider_env, auth_type, logger_mock):
async def _run_create(self, monkeypatch, provider_env, auth_type, caplog):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
add_mcp_server,
)
@ -3881,33 +3886,30 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap:
mock_manager.reload_servers_from_database = AsyncMock()
with (
patch(
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(
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(
patch( # test-quality-ok: endpoint reads the global MCP manager
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.verbose_proxy_logger",
logger_mock,
),
):
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"
),
)
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(
@ -3920,31 +3922,28 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap:
],
)
async def test_create_warns_under_a_provider_that_captures_nothing(
self, monkeypatch, provider_env, expected_fragment
self, monkeypatch, caplog, provider_env, expected_fragment
):
logger_mock = MagicMock()
await self._run_create(monkeypatch, provider_env, MCPAuth.oauth2_id_jag, logger_mock)
warnings = self._id_jag_warnings(logger_mock)
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):
logger_mock = MagicMock()
await self._run_create(monkeypatch, {"GENERIC_CLIENT_ID": "cid"}, MCPAuth.oauth2_id_jag, logger_mock)
assert self._id_jag_warnings(logger_mock) == []
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):
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."""
logger_mock = MagicMock()
await self._run_create(monkeypatch, {"GOOGLE_CLIENT_ID": "cid"}, MCPAuth.api_key, logger_mock)
assert self._id_jag_warnings(logger_mock) == []
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):
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,
@ -3956,42 +3955,37 @@ class TestIdJagRegistrationWarnsAboutTheSSOGap:
mock_manager = MagicMock()
mock_manager.update_server = AsyncMock()
mock_manager.reload_servers_from_database = AsyncMock()
logger_mock = MagicMock()
with (
patch(
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(
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(
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(
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(
patch( # test-quality-ok: endpoint reads the global MCP manager
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.verbose_proxy_logger",
logger_mock,
),
):
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"
),
)
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(logger_mock)
warnings = self._id_jag_warnings(caplog)
assert len(warnings) == 1
assert "google" in str(warnings[0])

View file

@ -1,5 +1,6 @@
import asyncio
import json
import logging
import os
from contextlib import ExitStack, asynccontextmanager
from types import SimpleNamespace
@ -8191,20 +8192,24 @@ async def _render_debug_page(provider_env, id_jag_registered, force_inert=False)
stack = [
patch.dict(os.environ, provider_env, clear=False),
patch("litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic),
patch.object(GoogleSSOHandler, "get_google_callback_response", side_effect=fake_google),
patch(
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", {}),
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)),
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(
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),
)
@ -8222,15 +8227,14 @@ async def _render_debug_page(provider_env, id_jag_registered, force_inert=False)
@pytest.mark.asyncio
async def test_debug_page_logs_the_capture_gap_but_never_renders_it():
warning_mock = MagicMock()
with patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger.warning", warning_mock):
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)
warning_mock.assert_called_once()
warning_args = " ".join(str(arg) for arg in warning_mock.call_args.args)
assert "google" in warning_args
assert "GENERIC_CLIENT_ID" in warning_args
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
@ -8270,23 +8274,20 @@ async def test_debug_page_is_byte_identical_when_no_id_jag_server_is_registered(
@pytest.mark.asyncio
async def test_debug_page_survives_a_store_outage():
async def test_debug_page_survives_a_store_outage(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
warning_mock = MagicMock()
with (
patch.dict(os.environ, {"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID}, clear=False),
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(side_effect=Exception("db down")),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger.warning", warning_mock),
):
assert await warn_if_id_jag_capture_gap() is None
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
assert (
await warn_if_id_jag_capture_gap(
retention_enabled=AsyncMock(side_effect=Exception("db down"))
)
is None
)
warning_mock.assert_not_called()
assert _id_jag_gap_warnings(caplog) == []
async def _render_legacy_login_page(env_overrides, general_settings):
@ -8803,8 +8804,12 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
assert response.status_code == 200
def _id_jag_gap_warnings(logger_mock) -> list:
return [call for call in logger_mock.warning.call_args_list if "oauth2_id_jag" in str(call)]
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
@ -8817,7 +8822,7 @@ def _id_jag_gap_warnings(logger_mock) -> list:
],
)
async def test_uncaptured_assertion_warns_when_an_id_jag_server_is_registered(
monkeypatch, provider_env, expected_fragment
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."""
@ -8830,23 +8835,16 @@ async def test_uncaptured_assertion_warns_when_an_id_jag_server_is_registered(
for name, value in provider_env.items():
monkeypatch.setenv(name, value)
logger_mock = MagicMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=True),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
await warn_if_id_jag_assertion_uncaptured(None)
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(logger_mock)
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):
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 (
@ -8857,23 +8855,16 @@ async def test_generic_provider_that_returned_no_id_token_still_warns(monkeypatc
monkeypatch.delenv(name, raising=False)
monkeypatch.setenv("GENERIC_CLIENT_ID", "cid")
logger_mock = MagicMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=True),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
await warn_if_id_jag_assertion_uncaptured(None)
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(logger_mock)
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):
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,
)
@ -8886,22 +8877,15 @@ async def test_no_warning_when_the_assertion_was_captured(monkeypatch):
assert assertion is not None
retention_mock = AsyncMock(return_value=True)
logger_mock = MagicMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
retention_mock,
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
await warn_if_id_jag_assertion_uncaptured(assertion)
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(logger_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):
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 (
@ -8909,103 +8893,98 @@ async def test_no_warning_when_no_id_jag_server_is_registered(monkeypatch):
)
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
logger_mock = MagicMock()
with (
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(return_value=False),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
await warn_if_id_jag_assertion_uncaptured(None)
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(logger_mock) == []
assert _id_jag_gap_warnings(caplog) == []
@pytest.mark.asyncio
async def test_store_outage_does_not_break_the_login(monkeypatch):
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 (
patch(
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
AsyncMock(side_effect=Exception("db down")),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", MagicMock()),
):
await warn_if_id_jag_assertion_uncaptured(None)
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):
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 = {}
logger_mock = MagicMock()
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(
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(
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(
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(
patch( # test-quality-ok: endpoint test stubs assertion persistence
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
patch(
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),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
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,
)
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(logger_mock)
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):
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,
@ -9021,49 +9000,48 @@ async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch):
user_info.models = []
user_info.teams = []
logger_mock = MagicMock()
with (
patch(
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(
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=[]),
),
patch(
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(
patch( # test-quality-ok: endpoint test stubs assertion persistence
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
AsyncMock(),
),
patch(
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),
),
patch("litellm.proxy.management_endpoints.ui_sso.verbose_proxy_logger", logger_mock),
):
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,
)
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(logger_mock)
warnings = _id_jag_gap_warnings(caplog)
assert len(warnings) == 1
assert "microsoft" in str(warnings[0])