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 387843ee5b2..bb3f1bece75 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 @@ -357,6 +357,7 @@ class MCPRequestHandler: # Inline imports avoid a circular dependency: mcp_server_manager imports # from this module. from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, global_mcp_server_manager, ) from litellm.types.mcp import MCPAuth @@ -382,7 +383,18 @@ class MCPRequestHandler: # fetches the upstream token automatically using stored credentials, # so allowing anonymous bypass would let any external caller invoke # tools authenticated as LiteLLM's service account. - if server.has_client_credentials: + # + # Resolve the flow rather than reading has_client_credentials directly: + # this is a security gate, and a legacy row whose oauth2_flow was never + # stamped still carries the M2M credential shape (client_id/secret + + # token_url, no authorization_url). Treating an unstamped-but-M2M-shaped + # row as non-M2M here would reopen the anonymous bypass the explicit + # column no longer closes on its own. Shares the one resolution helper + # with the egress backstop and the anonymous-delegate allowlist; all fail + # closed on the ambiguous shape and are removed together once no null rows + # remain. A pure-PKCE delegate server (no stored credentials) resolves to a + # non-M2M flow and keeps its bypass. + if MCPServerManager.effective_oauth2_flow(server) == "client_credentials": return False return True diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0d28d4d26c4..61ba49729cd 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -581,6 +581,21 @@ def _create_elicitation_callback(): class MCPServerManager: _STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$") + @staticmethod + def _explicit_oauth2_flow( + oauth2_flow: Optional[str], + ) -> Optional[Literal["client_credentials", "authorization_code"]]: + """DB rows persist their flow (write-time stamps plus the startup backfill) and + config servers must declare it (validated at load), so both builds read the + value verbatim: unknown or null resolves to None, which + ``needs_user_oauth_token`` already treats as interactive. Field-shape inference + survives only in the request-time security helpers (``effective_oauth2_flow`` / + ``resolve_oauth2_flow_for_request``). + """ + if oauth2_flow in ("client_credentials", "authorization_code"): + return cast(Literal["client_credentials", "authorization_code"], oauth2_flow) + return None + @staticmethod def _resolve_oauth2_flow( *, @@ -591,11 +606,15 @@ class MCPServerManager: client_id: Optional[str], client_secret: Optional[str], ) -> Optional[Literal["client_credentials", "authorization_code"]]: - """Infer oauth2_flow for legacy records that omit the field. + """Infer oauth2_flow from field shape when the value is omitted. - DB rows created before oauth2_flow support may have OAuth2 client - credentials + token_url but a null oauth2_flow. Treat these as M2M, - unless authorization_url is present (interactive OAuth). + Not called directly by security sites; they go through ``effective_oauth2_flow`` + (boolean/enum decisions) or ``resolve_oauth2_flow_for_request`` (the egress object + backstop), which are the single choke points for request-time resolution. DB rows + are stamped at write time and by the startup backfill, config servers must declare + oauth2_flow (validated at load), and both builds read the value verbatim via + ``_explicit_oauth2_flow``. Delete this whole request-time layer once the backstop + warning stays silent in production. """ if oauth2_flow in ("client_credentials", "authorization_code"): return cast(Literal["client_credentials", "authorization_code"], oauth2_flow) @@ -610,6 +629,51 @@ class MCPServerManager: return "client_credentials" return None + @staticmethod + def effective_oauth2_flow(server: "MCPServer") -> Optional[Literal["client_credentials", "authorization_code"]]: + """The oauth2_flow a security decision must use for ``server`` this request. + + Column-first, shape-fallback: a stamped row returns its explicit value; an + unstamped (null) row whose fields carry the M2M shape resolves to + ``client_credentials`` so it is treated as M2M and fails closed. Every + security-sensitive reader (anonymous-delegate allowlist and gate, egress flow + resolution) goes through this one helper rather than reading the bare + ``has_client_credentials`` column, which is unreliable for null rows. + """ + return MCPServerManager._resolve_oauth2_flow( + auth_type=server.auth_type, + oauth2_flow=server.oauth2_flow, + token_url=server.token_url, + authorization_url=server.authorization_url, + client_id=server.client_id, + client_secret=server.client_secret, + ) + + @staticmethod + def resolve_oauth2_flow_for_request(server: "MCPServer") -> "MCPServer": + """Return ``server`` with its effective oauth2_flow applied, for egress paths. + + A stamped row is returned unchanged (its effective flow equals the stored value). + An unstamped M2M-shape row is returned as a per-request copy carrying + ``oauth2_flow=client_credentials`` so downstream ``has_client_credentials`` / + ``needs_user_oauth_token`` compute correctly and the stored client credentials are + used instead of forwarding the caller's Authorization. Use this at every point that + resolves an allowed server id into an ``MCPServer`` for a tool call or listing. + """ + effective = MCPServerManager.effective_oauth2_flow(server) + if effective is None or effective == server.oauth2_flow: + return server + verbose_logger.warning( + "MCP server %s has no persisted oauth2_flow but matches the %s shape; using the " + "inferred flow for this request. The startup backfill leaves this ambiguous M2M " + "shape unstamped on purpose, so it will NOT self-heal: set oauth2_flow explicitly " + "in the dashboard or via PUT /v1/mcp/server (client_credentials for M2M, or " + "authorization_code after an interactive sign-in).", + server.server_id, + effective, + ) + return server.model_copy(update={"oauth2_flow": effective}) + @staticmethod def _obo_needs_endpoint_discovery( auth_type: Optional[MCPAuthType], @@ -842,6 +906,20 @@ class MCPServerManager: mcp_oauth_metadata.registration_url if mcp_oauth_metadata else None ) + config_oauth2_flow = server_config.get("oauth2_flow", None) + if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in ( + "client_credentials", + "authorization_code", + ): + raise ValueError( + f"Invalid config for MCP server '{server_name or server_id}': auth_type oauth2 " + f"requires an explicit oauth2_flow (got {config_oauth2_flow!r}). Set " + "oauth2_flow: client_credentials for machine-to-machine servers (the proxy mints " + "a shared token at token_url using client_id/client_secret, no user interaction) " + "or oauth2_flow: authorization_code for interactive servers (per-user tokens via " + "browser sign-in, including delegate_auth_to_upstream)." + ) + new_server = MCPServer( server_id=server_id, name=name_for_prefix, @@ -855,14 +933,7 @@ class MCPServerManager: # oauth specific fields client_id=server_config.get("client_id", None), client_secret=server_config.get("client_secret", None), - oauth2_flow=self._resolve_oauth2_flow( - auth_type=auth_type, - oauth2_flow=server_config.get("oauth2_flow", None), - token_url=resolved_token_url, - authorization_url=resolved_authorization_url, - client_id=server_config.get("client_id", None), - client_secret=server_config.get("client_secret", None), - ), + oauth2_flow=self._explicit_oauth2_flow(config_oauth2_flow), scopes=resolved_scopes, authorization_url=resolved_authorization_url, token_url=resolved_token_url, @@ -1240,15 +1311,7 @@ class MCPServerManager: env_vars=env_vars_list, client_id=client_id_value or getattr(mcp_server, "client_id", None), client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), - oauth2_flow=self._resolve_oauth2_flow( - auth_type=auth_type, - oauth2_flow=getattr(mcp_server, "oauth2_flow", None), - token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None), - authorization_url=mcp_server.authorization_url - or getattr(mcp_oauth_metadata, "authorization_url", None), - client_id=client_id_value or getattr(mcp_server, "client_id", None), - client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), - ), + oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)), scopes=resolved_scopes, authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None), token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None), @@ -1556,8 +1619,11 @@ class MCPServerManager: and getattr(server, "delegate_auth_to_upstream", False) is True # M2M servers must not be exposed anonymously: an # unauthenticated caller would get LiteLLM to proxy tool - # calls using its stored client_credentials. - and not server.has_client_credentials + # calls using its stored client_credentials. Resolve the flow + # rather than reading has_client_credentials so an unstamped + # M2M-shape row (null column, verbatim-read as non-M2M) still + # fails closed here, matching the anonymous-delegate auth gate. + and MCPServerManager.effective_oauth2_flow(server) != "client_credentials" ] combined_servers.update(delegate_server_ids) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d978771f433..e3812522ded 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1427,18 +1427,8 @@ if MCP_AVAILABLE: for allowed_mcp_server_id in allowed_mcp_server_ids: mcp_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) if mcp_server is not None: - # Apply oauth2_flow resolution for legacy DB rows where it may be NULL - resolved_flow = MCPServerManager._resolve_oauth2_flow( - auth_type=mcp_server.auth_type, - oauth2_flow=mcp_server.oauth2_flow, - token_url=mcp_server.token_url, - authorization_url=mcp_server.authorization_url, - client_id=mcp_server.client_id, - client_secret=mcp_server.client_secret, - ) - if resolved_flow and resolved_flow != mcp_server.oauth2_flow: - # Create a new instance with the resolved flow for this request - mcp_server = mcp_server.model_copy(update={"oauth2_flow": resolved_flow}) + # Apply the request-time oauth2_flow backstop for legacy null rows. + mcp_server = MCPServerManager.resolve_oauth2_flow_for_request(mcp_server) allowed_mcp_servers.append(mcp_server) if mcp_servers is not None: @@ -2800,6 +2790,9 @@ if MCP_AVAILABLE: for allowed_mcp_server_id in allowed_mcp_server_ids: allowed_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) if allowed_server is not None: + # Same request-time oauth2_flow backstop the listing path applies, + # so a null-flow M2M-shape row is treated as M2M on tool calls too. + allowed_server = MCPServerManager.resolve_oauth2_flow_for_request(allowed_server) allowed_mcp_servers.append(allowed_server) allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( 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 3607d448aad..06858320cbf 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 @@ -2146,6 +2146,104 @@ class TestMCPDelegateAuthToUpstream: assert exc_info.value.status_code == 401 mock_auth.assert_called_once() + async def test_delegate_ignored_for_unstamped_m2m_shaped_server(self): + """ + oauth2 + delegate + oauth2_flow=None but the M2M credential shape + (client_id/secret + token_url, no authorization_url) → bypass must NOT + fire. A legacy row that was never stamped still resolves to + client_credentials by shape, and reading the bare column here would + reopen the anonymous bypass to a server that runs upstream as LiteLLM's + service account. Fails closed like the client_credentials case above. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/legacy_m2m_server", + "headers": [], + } + + legacy_m2m_server = MCPServer( + server_id="legacy-m2m-id", + name="legacy_m2m_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow=None, + client_id="cid", + client_secret="csecret", + token_url="https://idp.example.com/token", + ) + assert legacy_m2m_server.has_client_credentials is False + + async def mock_auth_raises(*_args, **_kwargs): + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = legacy_m2m_server + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + mock_auth.assert_called_once() + + async def test_delegate_bypass_for_pure_pkce_server(self): + """ + oauth2 + delegate + oauth2_flow=None and NO stored client credentials + (pure PKCE, the common delegate case) → bypass must still fire. The + shape resolves to a non-M2M flow, so the security gate leaves it alone; + the fail-closed rule targets the M2M shape specifically, not every + unstamped row. + """ + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/pkce_server", + "headers": [], + } + + pkce_server = MCPServer( + server_id="pkce-server-id", + name="pkce_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow=None, + ) + + async def mock_auth_raises(*_args, **_kwargs): + from fastapi import HTTPException + + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = pkce_server + auth, *_rest = await MCPRequestHandler.process_mcp_request(scope) + mock_auth.assert_not_called() + assert auth.api_key is None + async def test_delegate_bypass_for_internal_server(self): """ Delegate + oauth2 interactive servers bypass LiteLLM auth even when @@ -2234,6 +2332,56 @@ class TestMCPDelegateAuthToUpstream: assert "pkce-server" in result assert "m2m-server" not in result + async def test_get_allowed_servers_excludes_unstamped_m2m_shape_delegate(self): + """ + The anonymous allow-list must also exclude an M2M-shape delegate server whose + oauth2_flow was never stamped (null column, verbatim-read as non-M2M). Reading + the bare has_client_credentials here would surface it to anonymous callers; the + resolved-flow check fails closed on the shape, matching the auth gate. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + pkce_server = MCPServer( + server_id="pkce-server", + name="pkce_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=True, + ) + unstamped_m2m = MCPServer( + server_id="unstamped-m2m", + name="unstamped_m2m", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow=None, + client_id="cid", + client_secret="csecret", + token_url="https://idp.example.com/token", + ) + assert unstamped_m2m.has_client_credentials is False + manager.registry = { + pkce_server.server_id: pkce_server, + unstamped_m2m.server_id: unstamped_m2m, + } + + with patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + result = await manager.get_allowed_mcp_servers(None) + + assert "pkce-server" in result + assert "unstamped-m2m" not in result + async def test_get_allowed_servers_includes_internal_delegate(self): """ Internal-only (available_on_public_internet=False) delegate servers diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 44f1d105093..1f44160aef4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6592,3 +6592,65 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_ assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == [] prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): + """ + Finding 3 regression: the call_mcp_tool path must apply the same request-time + oauth2_flow backstop the listing path does. A legacy DB row with oauth2_flow=NULL + but the M2M credential shape must reach execute_mcp_tool resolved to + client_credentials, or the caller's Authorization would be forwarded to an M2M + upstream on tool execution during a backfill gap (the list path was covered, the + call path was not). + """ + try: + from litellm.proxy._experimental.mcp_server.server import call_mcp_tool + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp import MCPAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="sk-1234", user_id="test-user") + + legacy_server = MCPServer( + server_id="legacy-m2m-id", + name="legacy_m2m", + alias="legacy_m2m", + server_name="legacy_m2m", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=None, # legacy: unstamped + token_url="https://oauth.example.com/token", + client_id="client-id", + client_secret="client-secret", + ) + assert legacy_server.has_client_credentials is False + + captured_servers = {} + + async def capture_execute(*args, **kwargs): + captured_servers["allowed"] = kwargs.get("allowed_mcp_servers") + return MagicMock(name="call_tool_result") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + side_effect=capture_execute, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers: allowed_mcp_servers), + ), + ): + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) + + await call_mcp_tool(name="legacy_m2m-tool", arguments={}, user_api_key_auth=user_auth) + + resolved = captured_servers["allowed"] + assert resolved and resolved[0].oauth2_flow == "client_credentials" + assert resolved[0].has_client_credentials is True diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 358c0409db4..306c74c0d83 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -293,6 +293,86 @@ class TestMCPServerManager: assert server.alias == "friendly_alias" assert server.server_name == "validserver" + def _oauth2_config(self, **overrides): + base = { + "url": "https://example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "token_url": "https://idp.example.com/token", + "client_id": "cid", + "client_secret": "csec", + } + base.update(overrides) + return {"m2mserver": base} + + @pytest.mark.asyncio + async def test_load_servers_from_config_requires_oauth2_flow(self): + """auth_type oauth2 without an explicit oauth2_flow is a config error: the + credential shape is ambiguous (a DCR interactive server looks identical to M2M), + so the config must assert the flow instead of the proxy guessing it.""" + + manager = MCPServerManager() + + with ( + patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), + pytest.raises(ValueError) as exc_info, + ): + await manager.load_servers_from_config(self._oauth2_config()) + + assert "oauth2_flow: client_credentials" in str(exc_info.value) + assert "oauth2_flow: authorization_code" in str(exc_info.value) + + @pytest.mark.asyncio + async def test_load_servers_from_config_rejects_unknown_oauth2_flow(self): + manager = MCPServerManager() + + with ( + patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), + pytest.raises(ValueError) as exc_info, + ): + await manager.load_servers_from_config(self._oauth2_config(oauth2_flow="m2m")) + + assert "got 'm2m'" in str(exc_info.value) + + @pytest.mark.asyncio + async def test_load_servers_from_config_accepts_explicit_client_credentials(self): + manager = MCPServerManager() + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + await manager.load_servers_from_config(self._oauth2_config(oauth2_flow="client_credentials")) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.oauth2_flow == "client_credentials" + assert server.has_client_credentials is True + + @pytest.mark.asyncio + async def test_load_servers_from_config_accepts_explicit_authorization_code(self): + manager = MCPServerManager() + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + await manager.load_servers_from_config(self._oauth2_config(oauth2_flow="authorization_code")) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.oauth2_flow == "authorization_code" + assert server.needs_user_oauth_token is True + + @pytest.mark.asyncio + async def test_load_servers_from_config_non_oauth2_needs_no_flow(self): + manager = MCPServerManager() + config = { + "apiserver": { + "url": "https://example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.api_key, + "auth_value": "sk-upstream", + } + } + + await manager.load_servers_from_config(config) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.oauth2_flow is None + @pytest.mark.asyncio async def test_load_servers_from_config_coerces_cost_string_to_float(self): """YAML 1.1 parses `7e-05` as a string; ingest must coerce it to float.""" @@ -1637,6 +1717,7 @@ class TestMCPServerManager: "url": "https://example.com/mcp", "transport": MCPTransport.http, "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", "scopes": ["config"], "authorization_url": "https://config.example.com/auth", } @@ -1700,6 +1781,7 @@ class TestMCPServerManager: "url": "https://example.com/mcp", "transport": MCPTransport.http, "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", "scopes": ["config"], "authorization_url": "https://config.example.com/auth", } @@ -6076,3 +6158,127 @@ async def test_aggregate_list_still_absorbs_step_up_challenged_server(): result = await manager.list_tools() assert [t.name for t in result] == ["good-do_thing"] + + +class TestDbBuildReadsOauth2FlowColumnVerbatim: + """The DB build must not re-infer the flow from field shape: rows are stamped at + write time and by the startup backfill, and a DCR-registered interactive server + has the exact M2M shape (client creds + token_url, no persisted authorization_url) + whenever discovery is unavailable. Inference survives only for config-loaded + servers and the request-time backstop in _get_allowed_mcp_servers.""" + + def _row(self, oauth2_flow): + return LiteLLM_MCPServerTable( + server_id="flow-column-row", + alias="flow_column_row", + description="", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=oauth2_flow, + token_url="https://idp.example.com/token", + credentials={"client_id": "cid", "client_secret": "csec"}, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + @pytest.mark.asyncio + async def test_null_flow_m2m_shape_row_is_not_inferred_m2m(self): + manager = MCPServerManager() + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + built = await manager.build_mcp_server_from_table(self._row(None), credentials_are_encrypted=False) + + assert built.oauth2_flow is None + assert built.has_client_credentials is False + assert built.needs_user_oauth_token is True + + @pytest.mark.asyncio + async def test_explicit_flow_column_is_read_verbatim(self): + manager = MCPServerManager() + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + built = await manager.build_mcp_server_from_table( + self._row("client_credentials"), credentials_are_encrypted=False + ) + + assert built.oauth2_flow == "client_credentials" + assert built.has_client_credentials is True + assert built.needs_user_oauth_token is False + + @pytest.mark.asyncio + async def test_authorization_code_flow_column_is_read_verbatim(self): + manager = MCPServerManager() + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + built = await manager.build_mcp_server_from_table( + self._row("authorization_code"), credentials_are_encrypted=False + ) + + assert built.oauth2_flow == "authorization_code" + assert built.has_client_credentials is False + assert built.needs_user_oauth_token is True + + +class TestRequestTimeOauth2FlowBackstop: + """The single request-time resolution helpers every security site shares: + effective_oauth2_flow (the enum/boolean decision) and + resolve_oauth2_flow_for_request (the egress object copy).""" + + def _oauth2_server(self, **overrides): + base = dict( + server_id="flow-server", + name="flow_server", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + base.update(overrides) + return MCPServer(**base) + + def test_effective_flow_stamped_values_returned_verbatim(self): + assert ( + MCPServerManager.effective_oauth2_flow(self._oauth2_server(oauth2_flow="client_credentials")) + == "client_credentials" + ) + assert ( + MCPServerManager.effective_oauth2_flow(self._oauth2_server(oauth2_flow="authorization_code")) + == "authorization_code" + ) + + def test_effective_flow_null_m2m_shape_resolves_client_credentials(self): + server = self._oauth2_server( + oauth2_flow=None, + client_id="cid", + client_secret="csecret", + token_url="https://idp.example.com/token", + ) + assert MCPServerManager.effective_oauth2_flow(server) == "client_credentials" + + def test_effective_flow_null_pure_pkce_resolves_none(self): + assert MCPServerManager.effective_oauth2_flow(self._oauth2_server(oauth2_flow=None)) is None + + def test_resolve_for_request_stamped_row_is_unchanged_identity(self): + server = self._oauth2_server(oauth2_flow="client_credentials") + assert MCPServerManager.resolve_oauth2_flow_for_request(server) is server + + def test_resolve_for_request_null_pure_pkce_is_unchanged_identity(self): + server = self._oauth2_server(oauth2_flow=None) + assert MCPServerManager.resolve_oauth2_flow_for_request(server) is server + + def test_resolve_for_request_null_m2m_shape_copies_client_credentials(self, caplog): + import logging + + server = self._oauth2_server( + oauth2_flow=None, + client_id="cid", + client_secret="csecret", + token_url="https://idp.example.com/token", + ) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + resolved = MCPServerManager.resolve_oauth2_flow_for_request(server) + + assert resolved is not server + assert resolved.oauth2_flow == "client_credentials" + assert server.oauth2_flow is None # original untouched + # Finding 2: the warning must NOT promise the backfill will stamp this row. + joined = " ".join(caplog.messages) + assert "no persisted oauth2_flow" in joined + assert "next proxy boot" not in joined + assert "will NOT self-heal" in joined