diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d5f4fa2b256..5f9384c7eea 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -254,6 +254,7 @@ _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on")) _OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS: Final = (0.05, 0.15) _OAUTH_DISCOVERY_RETRY_BASE_SECONDS: Final = 30.0 _OAUTH_DISCOVERY_RETRY_MAX_SECONDS: Final = 900.0 +_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS: Final = 300.0 def _oauth_discovery_now() -> float: @@ -1898,6 +1899,10 @@ class MCPServerManager: slot: Final = self._oauth_discovery_slot(server_id) return slot is not None and slot.generation == generation + def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None: + if self._oauth_discovery_slot_is_current(server_id, generation): + self._remove_oauth_discovery_slot(server_id) + def _publish_resolved_oauth_server( self, server: MCPServer, @@ -1910,6 +1915,12 @@ class MCPServerManager: elif server.server_id in self.config_mcp_servers: self.config_mcp_servers[server.server_id] = server else: + asyncio.get_running_loop().call_later( + _OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS, + self._expire_temporary_oauth_discovery, + server.server_id, + generation, + ) return server self._remove_oauth_discovery_slot(server.server_id) return server 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 ecf78359191..f0998cff583 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 @@ -12758,3 +12758,39 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager._set_oauth_discovery_deferred(original.server_id, True) assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None assert manager.registry[original.server_id] is replacement + + +@pytest.mark.asyncio +async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + manager._set_oauth_discovery_deferred(server.server_id, True) + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + assert manager._oauth_discovery_slot(server.server_id) is not None + loop: Final = asyncio.get_running_loop() + expired: Final = loop.create_future() + with patch.object(loop, "time", return_value=loop.time() + 301): + loop.call_later(0, expired.set_result, None) + await expired + assert resolved.authorization_url == server.authorization_url + assert manager._oauth_discovery_slot(server.server_id) is None + + +def test_old_temporary_discovery_expiry_preserves_replacement() -> None: + manager: Final = MCPServerManager() + manager._set_oauth_discovery_deferred("reused-session", True) + old_slot: Final = manager._oauth_discovery_slot("reused-session") + assert old_slot is not None + manager._set_oauth_discovery_deferred("reused-session", True) + replacement: Final = manager._oauth_discovery_slot("reused-session") + manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation) + assert manager._oauth_discovery_slot("reused-session") is replacement + assert replacement is not None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None