From ddae6eac6bc8526a40220df816c72791c96cb8da Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 11:20:21 -0700 Subject: [PATCH] fix(mcp): judge EMA retention only against the config and DB authorities, never the registry snapshot --- .../sso_assertion_store.py | 15 ++++---- .../test_sso_assertion_store.py | 38 +++++++++++++------ 2 files changed, 35 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py index 73bb87dcc33..e0927cc4f64 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -86,19 +86,20 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden async def ema_assertion_retention_enabled() -> bool: """Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only - retains bearer material while an EMA upstream exists to spend it on. The local registry is - the fast path; when it has none, the DB is consulted as the authoritative source, because - the registry is a per-process snapshot while the assertion write targets the shared DB. A - server added on another pod (or not yet loaded during startup) must still enable retention - here, or the drop only surfaces later as an unexplained challenge at the EMA upstream.""" + retains bearer material while an EMA upstream exists to spend it on. Judged against the two + configuration authorities: the pod-local config declaration and the shared DB row. The + in-memory registry is deliberately not consulted in either direction; it is a per-process + snapshot of the DB state that can be stale both ways (a server added on another pod would + silently drop the write, one removed on another pod would keep retaining bearer material), + and a gate guarding a shared-DB write must judge against that storage's authority.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle global_mcp_server_manager, ) from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global - servers = global_mcp_server_manager.get_registry().values() - if any(server.auth_type == MCPAuth.oauth2_id_jag for server in servers): + config_servers = global_mcp_server_manager.config_mcp_servers.values() + if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers): return True if prisma_client is None: return False diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py index 040bcbc1594..a3f46a49ba9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -136,12 +136,12 @@ async def test_retention_gate_requires_an_id_jag_server(): patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)), ): - manager.get_registry.return_value = { + manager.config_mcp_servers = { "s1": _server_with_auth(MCPAuth.oauth2), "s2": _server_with_auth(None), } assert await ema_assertion_retention_enabled() is False - manager.get_registry.return_value = { + manager.config_mcp_servers = { "s1": _server_with_auth(MCPAuth.oauth2), "s2": _server_with_auth(MCPAuth.oauth2_id_jag), } @@ -149,12 +149,11 @@ async def test_retention_gate_requires_an_id_jag_server(): @pytest.mark.asyncio -async def test_retention_gate_falls_back_to_db_when_registry_is_cold(): - """A server added on another pod (or before this pod's DB load) is invisible to the local - registry; the gate must still enable retention off the authoritative DB row, and read - False only when neither the registry nor the DB knows an id_jag server.""" +async def test_retention_gate_reads_the_db_when_config_declares_no_id_jag_server(): + """A DB-backed server added on another pod (or before this pod's DB load) must still enable + retention off the authoritative DB row; False only when neither authority knows one.""" with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager: - manager.get_registry.return_value = {"s1": _server_with_auth(MCPAuth.oauth2)} + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)} db_backed = _make_prisma({}, db_has_id_jag_server=True) with patch("litellm.proxy.proxy_server.prisma_client", db_backed): assert await ema_assertion_retention_enabled() is True @@ -165,6 +164,23 @@ async def test_retention_gate_falls_back_to_db_when_registry_is_cold(): assert await ema_assertion_retention_enabled() is False +@pytest.mark.asyncio +async def test_retention_gate_never_consults_the_registry_snapshot(): + """The registry is a per-process snapshot of DB state, stale in either direction: trusting + it positively would keep retaining bearer material after the last EMA server was removed on + another pod, trusting it negatively would drop writes for one added elsewhere. The gate must + judge only the config declaration and the DB row, so a stale snapshot listing an id_jag + server changes nothing.""" + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)), + ): + manager.config_mcp_servers = {} + manager.get_registry.return_value = {"stale": _server_with_auth(MCPAuth.oauth2_id_jag)} + assert await ema_assertion_retention_enabled() is False + manager.get_registry.assert_not_called() + + @pytest.mark.asyncio async def test_retain_persists_when_only_the_db_knows_the_id_jag_server(): stored = {} @@ -173,7 +189,7 @@ async def test_retain_persists_when_only_the_db_knows_the_id_jag_server(): patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, patch("litellm.proxy.proxy_server.prisma_client", prisma), ): - manager.get_registry.return_value = {} + manager.config_mcp_servers = {} await retain_sso_identity_assertion_for_ema( user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) ) @@ -247,7 +263,7 @@ async def test_retain_noop_when_no_id_jag_server(): patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, patch("litellm.proxy.proxy_server.prisma_client", prisma), ): - manager.get_registry.return_value = {"s1": _server_with_auth(MCPAuth.oauth2)} + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)} await retain_sso_identity_assertion_for_ema( user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) ) @@ -263,7 +279,7 @@ async def test_retain_persists_when_id_jag_server_registered(): patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, patch("litellm.proxy.proxy_server.prisma_client", prisma), ): - manager.get_registry.return_value = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} await retain_sso_identity_assertion_for_ema( user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) ) @@ -289,7 +305,7 @@ async def test_retain_swallows_store_failure(): patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, patch("litellm.proxy.proxy_server.prisma_client", prisma), ): - manager.get_registry.return_value = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} await retain_sso_identity_assertion_for_ema( user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) )