mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: ensure_async_clients_test error
This commit is contained in:
parent
100f617a8f
commit
a8092ab4c4
2 changed files with 76 additions and 98 deletions
|
|
@ -7,9 +7,11 @@ 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,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
|
@ -214,28 +216,28 @@ def create_tool_function(
|
|||
except (json.JSONDecodeError, TypeError):
|
||||
json_body = {"data": body_value}
|
||||
|
||||
# Make HTTP request
|
||||
async with httpx.AsyncClient() as client:
|
||||
if original_method == "get":
|
||||
response = await client.get(url, params=params, headers=headers)
|
||||
elif original_method == "post":
|
||||
response = await client.post(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
elif original_method == "put":
|
||||
response = await client.put(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
elif original_method == "delete":
|
||||
response = await client.delete(url, params=params, headers=headers)
|
||||
elif original_method == "patch":
|
||||
response = await client.patch(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
else:
|
||||
return f"Unsupported HTTP method: {original_method}"
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
|
||||
return response.text
|
||||
if original_method == "get":
|
||||
response = await client.get(url, params=params, headers=headers)
|
||||
elif original_method == "post":
|
||||
response = await client.post(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
elif original_method == "put":
|
||||
response = await client.put(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
elif original_method == "delete":
|
||||
response = await client.delete(url, params=params, headers=headers)
|
||||
elif original_method == "patch":
|
||||
response = await client.patch(
|
||||
url, params=params, json=json_body, headers=headers
|
||||
)
|
||||
else:
|
||||
return f"Unsupported HTTP method: {original_method}"
|
||||
|
||||
return response.text
|
||||
|
||||
return tool_function
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ This test suite ensures that:
|
|||
"""
|
||||
|
||||
import pytest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
|
|
@ -19,6 +20,19 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
|||
)
|
||||
|
||||
|
||||
GET_ASYNC_CLIENT_TARGET = (
|
||||
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client"
|
||||
)
|
||||
|
||||
|
||||
def _create_mock_client(method: str, response_text: str) -> AsyncMock:
|
||||
"""Utility to create a mocked async httpx client for the given method."""
|
||||
response = SimpleNamespace(text=response_text)
|
||||
client = AsyncMock()
|
||||
setattr(client, method, AsyncMock(return_value=response))
|
||||
return client
|
||||
|
||||
|
||||
class TestCreateToolFunction:
|
||||
"""Test create_tool_function with various parameter name edge cases."""
|
||||
|
||||
|
|
@ -48,18 +62,15 @@ class TestCreateToolFunction:
|
|||
assert func.__name__ == "tool_function"
|
||||
|
||||
# Test calling with original parameter name
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = '{"id": "123"}'
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", '{"id": "123"}')
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"repository-id": "test-repo"})
|
||||
assert result == '{"id": "123"}'
|
||||
|
||||
# Verify URL was constructed correctly
|
||||
call_args = mock_client.return_value.__aenter__.return_value.get.call_args
|
||||
call_args = async_client.get.call_args
|
||||
assert "repository-id" in str(call_args[0][0]) or "test-repo" in str(
|
||||
call_args[0][0]
|
||||
)
|
||||
|
|
@ -87,18 +98,15 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "verified"
|
||||
mock_client.return_value.__aenter__.return_value.post = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("post", "verified")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"2fa-code": "123456"})
|
||||
assert result == "verified"
|
||||
|
||||
# Verify query parameter was included
|
||||
call_args = mock_client.return_value.__aenter__.return_value.post.call_args
|
||||
call_args = async_client.post.call_args
|
||||
assert call_args[1]["params"]["2fa-code"] == "123456"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -124,17 +132,14 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "found"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "found")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"user.name": "john.doe"})
|
||||
assert result == "found"
|
||||
|
||||
call_args = mock_client.return_value.__aenter__.return_value.get.call_args
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["user.name"] == "john.doe"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -160,17 +165,14 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "[]"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "[]")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"$filter": "name eq 'test'"})
|
||||
assert result == "[]"
|
||||
|
||||
call_args = mock_client.return_value.__aenter__.return_value.get.call_args
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["$filter"] == "name eq 'test'"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -196,17 +198,14 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "items"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "items")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"class": "premium"})
|
||||
assert result == "items"
|
||||
|
||||
call_args = mock_client.return_value.__aenter__.return_value.get.call_args
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["class"] == "premium"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -244,12 +243,9 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "success"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "success")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(
|
||||
**{
|
||||
|
|
@ -286,17 +282,14 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "created"
|
||||
mock_client.return_value.__aenter__.return_value.post = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("post", "created")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"body": {"name": "test"}})
|
||||
assert result == "created"
|
||||
|
||||
call_args = mock_client.return_value.__aenter__.return_value.post.call_args
|
||||
call_args = async_client.post.call_args
|
||||
assert call_args[1]["json"] == {"name": "test"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -313,12 +306,9 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "ok"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "ok")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func()
|
||||
assert result == "ok"
|
||||
|
|
@ -349,20 +339,9 @@ class TestCreateToolFunction:
|
|||
|
||||
assert callable(func)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "success"
|
||||
|
||||
client_method = getattr(
|
||||
mock_client.return_value.__aenter__.return_value, method
|
||||
)
|
||||
client_method.return_value = mock_response
|
||||
client_method = AsyncMock(return_value=mock_response)
|
||||
setattr(
|
||||
mock_client.return_value.__aenter__.return_value,
|
||||
method,
|
||||
client_method,
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client(method, "success")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"repository-id": "test"})
|
||||
assert result == "success"
|
||||
|
|
@ -505,18 +484,15 @@ class TestPathSecurity:
|
|||
base_url="https://example.com",
|
||||
)
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.text = "dummy-response"
|
||||
mock_client.return_value.__aenter__.return_value.get = AsyncMock(
|
||||
return_value=mock_response
|
||||
)
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "dummy-response")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
response = await tool_function(**{"filename": "report 2024.json"})
|
||||
|
||||
assert response == "dummy-response"
|
||||
|
||||
# Verify URL was properly encoded
|
||||
call_args = mock_client.return_value.__aenter__.return_value.get.call_args
|
||||
call_args = async_client.get.call_args
|
||||
url = call_args[0][0]
|
||||
assert url == "https://example.com/files/report%202024.json"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue