diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 9312764aa89..bbbb8ad4a36 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -7,12 +7,12 @@ from pathlib import PurePosixPath from typing import Any, Dict, Optional from urllib.parse import quote -import httpx - from litellm._logging import verbose_logger +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) +from litellm.types.llms.custom_http import httpxSpecialProvider # Store the base URL and headers globally BASE_URL = "" @@ -92,26 +92,25 @@ async def _invoke_openapi_operation( except Exception: json_body = {"data": body_value} - async with httpx.AsyncClient() as client: - method_lower = method.lower() - if method_lower == "get": - response = await client.get(url, params=params, headers=headers) - elif method_lower == "post": - response = await client.post( - url, params=params, json=json_body, headers=headers - ) - elif method_lower == "put": - response = await client.put( - url, params=params, json=json_body, headers=headers - ) - elif method_lower == "delete": - response = await client.delete(url, params=params, headers=headers) - elif method_lower == "patch": - response = await client.patch( - url, params=params, json=json_body, headers=headers - ) - else: - return f"Unsupported HTTP method: {method}" + client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) + + method_lower = method.lower() + if method_lower == "get": + response = await client.get(url, params=params, headers=headers) + elif method_lower == "post": + response = await client.post( + url, params=params, json=json_body, headers=headers + ) + elif method_lower == "put": + response = await client.put(url, params=params, json=json_body, headers=headers) + elif method_lower == "delete": + response = await client.delete(url, params=params, headers=headers) + elif method_lower == "patch": + response = await client.patch( + url, params=params, json=json_body, headers=headers + ) + else: + return f"Unsupported HTTP method: {method}" return response.text @@ -251,7 +250,6 @@ async def tool_function({params_str}) -> str: # Execute the function code to create the actual function local_vars = { - "httpx": httpx, "headers": headers, "base_url": base_url, "path": path, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index b51ca27700e..71668d4af47 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -35,8 +35,8 @@ class _DummyAsyncClient: @pytest.mark.asyncio async def test_should_reject_path_traversal_inputs(monkeypatch): monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.httpx.AsyncClient", - _DummyAsyncClient, + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client", + lambda *_, **__: _DummyAsyncClient(), ) _DummyAsyncClient.last_instance = None @@ -66,8 +66,8 @@ async def test_should_reject_path_traversal_inputs(monkeypatch): @pytest.mark.asyncio async def test_should_encode_and_request_safe_path_parameters(monkeypatch): monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.httpx.AsyncClient", - _DummyAsyncClient, + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client", + lambda *_, **__: _DummyAsyncClient(), ) _DummyAsyncClient.last_instance = None