mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: add checking path param
This commit is contained in:
parent
7ba14a30ed
commit
6168e500a8
2 changed files with 133 additions and 1 deletions
|
|
@ -3,7 +3,9 @@ This module is used to generate MCP tools from OpenAPI specs.
|
|||
"""
|
||||
|
||||
import json
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Any, Dict, Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -17,6 +19,29 @@ BASE_URL = ""
|
|||
HEADERS: Dict[str, str] = {}
|
||||
|
||||
|
||||
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
|
||||
"""Ensure path params cannot introduce directory traversal."""
|
||||
if param_value is None:
|
||||
return ""
|
||||
|
||||
value_str = str(param_value)
|
||||
if value_str == "":
|
||||
return ""
|
||||
|
||||
normalized_value = value_str.replace("\\", "/")
|
||||
if "/" in normalized_value:
|
||||
raise ValueError(
|
||||
f"Path parameter '{param_name}' must not contain path separators"
|
||||
)
|
||||
|
||||
if any(part in {".", ".."} for part in PurePosixPath(normalized_value).parts):
|
||||
raise ValueError(
|
||||
f"Path parameter '{param_name}' cannot include '.' or '..' segments"
|
||||
)
|
||||
|
||||
return quote(value_str, safe="")
|
||||
|
||||
|
||||
def load_openapi_spec(filepath: str) -> Dict[str, Any]:
|
||||
"""Load OpenAPI specification from JSON file."""
|
||||
with open(filepath, "r") as f:
|
||||
|
|
@ -142,7 +167,12 @@ async def tool_function({params_str}) -> str:
|
|||
for param_name in path_param_names:
|
||||
param_value = locals().get(param_name, "")
|
||||
if param_value:
|
||||
url = url.replace("{{" + param_name + "}}", str(param_value))
|
||||
# url = url.replace("{{" + param_name + "}}", str(param_value))
|
||||
try:
|
||||
safe_value = _sanitize_path_parameter_value(param_value, param_name)
|
||||
except ValueError as exc:
|
||||
return "Invalid path parameter: " + str(exc)
|
||||
url = url.replace("{{" + param_name + "}}", safe_value)
|
||||
|
||||
# Build query params
|
||||
query_param_names = {query_params}
|
||||
|
|
@ -192,6 +222,7 @@ async def tool_function({params_str}) -> str:
|
|||
"base_url": base_url,
|
||||
"path": path,
|
||||
"method": method,
|
||||
"_sanitize_path_parameter_value": _sanitize_path_parameter_value,
|
||||
}
|
||||
exec(func_code, local_vars)
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,101 @@
|
|||
"""Tests for OpenAPI to MCP generator path handling."""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
create_tool_function,
|
||||
)
|
||||
|
||||
|
||||
class _DummyResponse:
|
||||
def __init__(self, text: str = "ok"):
|
||||
self.text = text
|
||||
|
||||
|
||||
class _DummyAsyncClient:
|
||||
"""Minimal async client stub that records requests."""
|
||||
|
||||
last_instance = None
|
||||
|
||||
def __init__(self):
|
||||
self.requests = []
|
||||
_DummyAsyncClient.last_instance = self
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
async def get(self, url, params=None, headers=None):
|
||||
self.requests.append(("get", url, params, headers))
|
||||
return _DummyResponse("dummy-response")
|
||||
|
||||
|
||||
@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,
|
||||
)
|
||||
_DummyAsyncClient.last_instance = None
|
||||
|
||||
operation = {
|
||||
"parameters": [
|
||||
{
|
||||
"name": "filename",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
tool_function = create_tool_function(
|
||||
path="/files/{filename}",
|
||||
method="GET",
|
||||
operation=operation,
|
||||
base_url="https://example.com",
|
||||
)
|
||||
|
||||
response = await tool_function(filename="../admin")
|
||||
|
||||
assert "Invalid path parameter" in response
|
||||
|
||||
|
||||
@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,
|
||||
)
|
||||
_DummyAsyncClient.last_instance = None
|
||||
|
||||
operation = {
|
||||
"parameters": [
|
||||
{
|
||||
"name": "filename",
|
||||
"in": "path",
|
||||
"required": True,
|
||||
"schema": {"type": "string"},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
tool_function = create_tool_function(
|
||||
path="/files/{filename}",
|
||||
method="GET",
|
||||
operation=operation,
|
||||
base_url="https://example.com",
|
||||
)
|
||||
|
||||
response = await tool_function(filename="report 2024.json")
|
||||
|
||||
assert response == "dummy-response"
|
||||
|
||||
dummy_client = _DummyAsyncClient.last_instance
|
||||
assert dummy_client is not None
|
||||
method, url, params, headers = dummy_client.requests[0]
|
||||
assert method == "get"
|
||||
assert url == "https://example.com/files/report%202024.json"
|
||||
assert params == {}
|
||||
Loading…
Add table
Reference in a new issue