From 1b28128b22132ac0b4038c3ebaec6da2414957fe Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 14 Jul 2026 09:57:17 -0700 Subject: [PATCH] refactor(mcp): bind org_id on session admission and clarify unreachable arm - _reload_admitted_user binds the user's org_id so the org-level MCP ceiling stays in force for a gateway session (a ceiling can only narrow; multi-org users are capped conservatively to their primary org) instead of being silently skipped - trim the NotSessionBearer arm comment to state it is simply unreachable --- .../mcp_server/auth/user_api_key_auth_mcp.py | 23 +++++---- .../auth/test_user_api_key_auth_mcp.py | 49 ++++++++----------- 2 files changed, 34 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 774c7181aa7..0ef7416b677 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -375,8 +375,7 @@ class MCPRequestHandler: route=request_route, ) elif ( - is_mcp_gateway_dcr_enabled() - and _is_aggregate_mcp_scope(request_route, mcp_servers) + _is_aggregate_mcp_scope(request_route, mcp_servers) and oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]) ): @@ -701,8 +700,8 @@ class MCPRequestHandler: case SessionBearerInvalid(): raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) case NotSessionBearer(): - # is_session_bearer_shaped gated entry, so a non-session bearer here means a - # session-shaped-but-empty value; fail closed with the same challenge. + # Unreachable: the arm is entered only for an is_session_bearer_shaped + # value. Kept for match exhaustiveness and fails closed regardless. raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) case _: assert_never(result) @@ -752,12 +751,15 @@ class MCPRequestHandler: The DCR client authenticates via SSO at the bridged authorize, which yields a user subject rather than a virtual key, so the envelope admits under the user's own - identity: the reloaded ``user_id`` and the user's own MCP object permission ride on the - returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` the key path uses then - computes which servers the user may reach, so the user's litellm MCP grants and access groups - gate the request exactly as a key's do. Only the user's OWN object permission is bound: a - ``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, so - team-inherited MCP grants for a user are a follow-up (they need a many-teams union + identity: the reloaded ``user_id``, the user's own MCP object permission, and the user's + ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` + the key path uses then computes which servers the user may reach, so the user's litellm MCP + grants and access groups gate the request exactly as a key's do. Binding ``org_id`` keeps the + org-level MCP ceiling in force for this admission rather than silently skipping it; a user's + primary organization is used, so a user who spans organizations is capped conservatively (the + ceiling can only narrow the result, never broaden it). Only the user's OWN object permission is + bound: a ``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, + so team-inherited MCP grants for a user are a follow-up (they need a many-teams union ``get_allowed_mcp_servers`` does not do off one auth object). The caller's centralized policy gate enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed. @@ -805,6 +807,7 @@ class MCPRequestHandler: return UserAPIKeyAuth( user_id=user_object.user_id, user_role=user_object.user_role, + org_id=user_object.organization_id, object_permission=object_permission, object_permission_id=user_object.object_permission_id, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index e31e0782722..33e78027872 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -5121,6 +5121,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=None, @@ -5163,6 +5164,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=object_permission, @@ -5240,7 +5242,7 @@ class TestMCPDcrBridgeDelegateAdmission: with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False})), + self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", organization_id=None, metadata={"scim_active": False})), ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() with pytest.raises(HTTPException) as exc_info: @@ -6087,10 +6089,9 @@ class TestGatewaySessionAdmission: """The aggregate /mcp session-bearer admission arm (mcp_gateway_dcr). A valid session token admits under the LIVE litellm user it references; an invalid/expired/refresh/foreign token fails closed with the aggregate invalid_token challenge; the arm fires ONLY at the - aggregate scope with the flag on, never for named servers or per-server flows.""" + aggregate scope, never for named servers or per-server flows.""" _MASTER_KEY = "sk-gateway-session-admission-master-key" - _FLAG = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.is_mcp_gateway_dcr_enabled" def _session_bearer(self, user_id="sso-user-42", client_id="llm_dcrc_abc"): from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( @@ -6122,10 +6123,11 @@ class TestGatewaySessionAdmission: @staticmethod @contextlib.contextmanager - def _patch_user_reload(*, user_id, active=True): + def _patch_user_reload(*, user_id, active=True, organization_id=None): get_user_object = AsyncMock( return_value=MagicMock( user_id=user_id, + organization_id=organization_id, metadata={"scim_active": active} if not active else {"scim_active": True}, user_role=None, object_permission=None, @@ -6139,10 +6141,24 @@ class TestGatewaySessionAdmission: ): yield get_user_object + async def test_session_admission_binds_org_id_so_the_org_ceiling_applies(self): + """The admitted auth carries the user's org_id, so get_allowed_mcp_servers keeps the + org-level MCP ceiling in force for a gateway session instead of skipping it.""" + token = self._access_token(user_id="org-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="org-user", organization_id="org-123"), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert auth_result.org_id == "org-123" + async def test_valid_session_admits_under_live_user_at_aggregate_scope(self): token = self._access_token(user_id="sso-user-42") with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", @@ -6166,7 +6182,6 @@ class TestGatewaySessionAdmission: mint, _refresh, principal, keys = self._session_bearer() token = mint(principal, keys, datetime(2020, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), ): with pytest.raises(HTTPException) as exc_info: @@ -6178,7 +6193,6 @@ class TestGatewaySessionAdmission: token = self._access_token() tampered = token[:-3] + ("aaa" if not token.endswith("aaa") else "bbb") with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), ): with pytest.raises(HTTPException) as exc_info: @@ -6191,7 +6205,6 @@ class TestGatewaySessionAdmission: _mint, refresh, principal, keys = self._session_bearer() refresh_token = refresh(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), ): with pytest.raises(HTTPException) as exc_info: @@ -6201,37 +6214,17 @@ class TestGatewaySessionAdmission: async def test_foreign_key_session_fails_closed(self): token = self._access_token() with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", "sk-a-totally-different-master-key"), ): with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(self._scope(token)) assert exc_info.value.status_code == 401 - async def test_arm_does_not_fire_when_flag_off(self): - """Flag off: a session-shaped bearer is treated as an ordinary bearer and hits the - oauth2 arm, which validates it as a litellm credential and fails it there (not the - session arm). Proven by user_api_key_auth being called, unlike the flag-on path.""" - token = self._access_token() - with ( - patch(self._FLAG, return_value=False), - patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", - new_callable=AsyncMock, - side_effect=ProxyException(message="bad key", type="auth_error", param="api_key", code=401), - ) as mock_auth, - ): - with pytest.raises((HTTPException, ProxyException)): - await MCPRequestHandler.process_mcp_request(self._scope(token)) - mock_auth.assert_called_once() - async def test_arm_does_not_fire_for_named_server(self): """A session-shaped bearer aimed at a named server (path scope) does not enter the aggregate arm; it is treated as an ordinary bearer on that server.""" token = self._access_token() with ( - patch(self._FLAG, return_value=True), patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",