From 4573ab326b5b4293261ab7795b056018c06d1fa7 Mon Sep 17 00:00:00 2001 From: hamzaq453 Date: Tue, 30 Dec 2025 14:04:15 +0500 Subject: [PATCH] refactor: remove to_safe_identifier mapping from OpenAPI MCP generator --- .../mcp_server/openapi_to_mcp_generator.py | 159 ++++----------- .../test_openapi_to_mcp_generator.py | 186 ++++++------------ 2 files changed, 91 insertions(+), 254 deletions(-) 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 dc5d0ce73af..e4969df131e 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -3,8 +3,6 @@ This module is used to generate MCP tools from OpenAPI specs. """ import json -import keyword -import re from typing import Any, Dict, Optional import httpx @@ -19,64 +17,6 @@ BASE_URL = "" HEADERS: Dict[str, str] = {} -def to_safe_identifier(name: str) -> str: - """ - Convert an OpenAPI parameter name to a safe Python identifier. - - This function ensures that any parameter name from an OpenAPI spec can be - used as a Python function parameter without causing syntax errors or - security issues. It handles: - - Hyphens, dots, and other special characters - - Leading digits - - Python keywords - - Special characters like $, @, etc. - - Args: - name: The original parameter name from the OpenAPI spec - - Returns: - A valid Python identifier that can be used in function signatures - - Examples: - >>> to_safe_identifier("repository-id") - 'repository_id' - >>> to_safe_identifier("2fa-code") - '_2fa_code' - >>> to_safe_identifier("user.name") - 'user_name' - >>> to_safe_identifier("$filter") - '_filter' - >>> to_safe_identifier("class") - 'class_' - """ - if not name: - return "_empty_" - - # Start with underscore if first char is not a letter - # Replace all non-alphanumeric chars (except underscore) with underscore - # Collapse multiple underscores - safe = re.sub(r'[^a-zA-Z0-9_]', '_', name) - safe = re.sub(r'_+', '_', safe) # Collapse multiple underscores - - # If starts with digit, prefix with underscore - if safe and safe[0].isdigit(): - safe = '_' + safe - - # If empty after sanitization, use a default - if not safe: - safe = '_param_' - - # If it's a Python keyword, append underscore - if keyword.iskeyword(safe): - safe = safe + '_' - - # Ensure it doesn't start with a digit (shouldn't happen after above, but double-check) - if safe and safe[0].isdigit(): - safe = '_' + safe - - return safe - - def load_openapi_spec(filepath: str) -> Dict[str, Any]: """Load OpenAPI specification from JSON file.""" with open(filepath, "r") as f: @@ -173,9 +113,8 @@ def create_tool_function( """Create a tool function for an OpenAPI operation. This function creates an async tool function that can be called with - keyword arguments. Parameter names from the OpenAPI spec are safely - mapped to valid Python identifiers to avoid syntax errors and security - issues. + keyword arguments. Parameter names from the OpenAPI spec are accessed + directly via **kwargs, avoiding syntax errors from invalid Python identifiers. Args: path: API endpoint path @@ -191,108 +130,80 @@ def create_tool_function( headers = {} path_params, query_params, body_params = extract_parameters(operation) - all_params = path_params + query_params + body_params - - # Create mapping from original parameter names to safe identifiers - # This allows us to accept kwargs with original names but use safe names internally - param_name_map: Dict[str, str] = {} - safe_to_original_map: Dict[str, str] = {} - - for orig_name in all_params: - safe_name = to_safe_identifier(orig_name) - # Handle collisions: if safe name already exists, append a counter - counter = 1 - original_safe = safe_name - while safe_name in safe_to_original_map: - safe_name = f"{original_safe}_{counter}" - counter += 1 - - param_name_map[orig_name] = safe_name - safe_to_original_map[safe_name] = orig_name - - # Store original parameter lists for use in the closure - original_path_params = path_params - original_query_params = query_params - original_body_params = body_params original_method = method.lower() async def tool_function(**kwargs: Any) -> str: """ Dynamically generated tool function. - + Accepts keyword arguments where keys are the original OpenAPI parameter names. - The function safely handles parameter names that aren't valid Python identifiers. + The function safely handles parameter names that aren't valid Python identifiers + by using **kwargs instead of named parameters. """ # Build URL from base_url and path url = base_url + path - + # Replace path parameters using original names from OpenAPI spec - for orig_param_name in original_path_params: - # Try to get value using original name first, then safe name - param_value = kwargs.get(orig_param_name, "") - if not param_value and orig_param_name in param_name_map: - safe_name = param_name_map[orig_param_name] - param_value = kwargs.get(safe_name, "") - + for param_name in path_params: + param_value = kwargs.get(param_name, "") if param_value: # Replace {param_name} or {{param_name}} in URL - url = url.replace("{" + orig_param_name + "}", str(param_value)) - url = url.replace("{{" + orig_param_name + "}}", str(param_value)) - + url = url.replace("{" + param_name + "}", str(param_value)) + url = url.replace("{{" + param_name + "}}", str(param_value)) + # Build query params using original parameter names params: Dict[str, Any] = {} - for orig_param_name in original_query_params: - # Try to get value using original name first, then safe name - param_value = kwargs.get(orig_param_name, "") - if not param_value and orig_param_name in param_name_map: - safe_name = param_name_map[orig_param_name] - param_value = kwargs.get(safe_name, "") - + for param_name in query_params: + param_value = kwargs.get(param_name, "") if param_value: # Use original parameter name in query string (as expected by API) - params[orig_param_name] = param_value - + params[param_name] = param_value + # Build request body json_body: Optional[Dict[str, Any]] = None - if original_body_params: + if body_params: # Try "body" first (most common), then check all body param names body_value = kwargs.get("body", {}) if not body_value: - for orig_param_name in original_body_params: - body_value = kwargs.get(orig_param_name, {}) + for param_name in body_params: + body_value = kwargs.get(param_name, {}) if body_value: break - # Also try safe name - if orig_param_name in param_name_map: - safe_name = param_name_map[orig_param_name] - body_value = kwargs.get(safe_name, {}) - if body_value: - break - + if isinstance(body_value, dict): json_body = body_value elif body_value: # If it's a string, try to parse as JSON try: - json_body = json.loads(body_value) if isinstance(body_value, str) else {"data": body_value} + json_body = ( + json.loads(body_value) + if isinstance(body_value, str) + else {"data": body_value} + ) 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) + 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) + 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) + 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 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 fa28dc34981..cb48a940b57 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 @@ -7,91 +7,16 @@ This test suite ensures that: 3. All edge cases (hyphens, dots, keywords, special chars) work correctly """ -import json import pytest from unittest.mock import AsyncMock, patch -from typing import Dict, Any from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - to_safe_identifier, create_tool_function, build_input_schema, extract_parameters, ) -class TestToSafeIdentifier: - """Test the to_safe_identifier function for various edge cases.""" - - def test_hyphen_in_name(self): - """Test parameter names with hyphens.""" - assert to_safe_identifier("repository-id") == "repository_id" - assert to_safe_identifier("user-name") == "user_name" - assert to_safe_identifier("api-key") == "api_key" - - def test_leading_digit(self): - """Test parameter names starting with digits.""" - assert to_safe_identifier("2fa-code") == "_2fa_code" - assert to_safe_identifier("123abc") == "_123abc" - assert to_safe_identifier("0test") == "_0test" - - def test_dots_in_name(self): - """Test parameter names with dots.""" - assert to_safe_identifier("user.name") == "user_name" - assert to_safe_identifier("config.value") == "config_value" - assert to_safe_identifier("api.v2") == "api_v2" - - def test_dollar_sign(self): - """Test parameter names with dollar signs (OData style).""" - assert to_safe_identifier("$filter") == "_filter" - assert to_safe_identifier("$context") == "_context" - assert to_safe_identifier("$select") == "_select" - - def test_python_keywords(self): - """Test Python keywords are handled.""" - assert to_safe_identifier("class") == "class_" - assert to_safe_identifier("from") == "from_" - assert to_safe_identifier("not") == "not_" - assert to_safe_identifier("def") == "def_" - assert to_safe_identifier("import") == "import_" - - def test_special_characters(self): - """Test various special characters.""" - assert to_safe_identifier("user@domain") == "user_domain" - assert to_safe_identifier("test#hash") == "test_hash" - assert to_safe_identifier("path/to/resource") == "path_to_resource" - assert to_safe_identifier("param+value") == "param_value" - - def test_multiple_special_chars(self): - """Test names with multiple special characters.""" - assert to_safe_identifier("user-name.email@domain") == "user_name_email_domain" - assert to_safe_identifier("$filter.value") == "_filter_value" - - def test_already_valid_identifier(self): - """Test that valid identifiers remain unchanged (except keywords).""" - assert to_safe_identifier("valid_name") == "valid_name" - assert to_safe_identifier("validName123") == "validName123" - assert to_safe_identifier("_private") == "_private" - - def test_empty_string(self): - """Test empty string handling.""" - assert to_safe_identifier("") == "_empty_" - - def test_only_special_chars(self): - """Test names that are only special characters.""" - result = to_safe_identifier("---") - assert result.startswith("_") - assert len(result) > 0 - - def test_collision_handling(self): - """Test that similar names produce different safe identifiers.""" - # These should produce different results - name1 = to_safe_identifier("user-name") - name2 = to_safe_identifier("user_name") - # They might be the same after sanitization, which is acceptable - # The important thing is they're both valid identifiers - - class TestCreateToolFunction: """Test create_tool_function with various parameter name edge cases.""" @@ -108,18 +33,18 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/repos/{repository-id}", method="get", operation=operation, base_url="https://api.example.com", ) - + # Should not raise SyntaxError assert callable(func) assert func.__name__ == "tool_function" - + # Test calling with original parameter name with patch("httpx.AsyncClient") as mock_client: mock_response = AsyncMock() @@ -127,13 +52,15 @@ class TestCreateToolFunction: mock_client.return_value.__aenter__.return_value.get = AsyncMock( return_value=mock_response ) - + 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 - assert "repository-id" in str(call_args[0][0]) or "test-repo" in str(call_args[0][0]) + assert "repository-id" in str(call_args[0][0]) or "test-repo" in str( + call_args[0][0] + ) @pytest.mark.asyncio async def test_leading_digit_parameter(self): @@ -148,26 +75,26 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/verify", method="post", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + 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 assert call_args[1]["params"]["2fa-code"] == "123456" @@ -185,26 +112,26 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/search", method="get", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func(**{"user.name": "john.doe"}) assert result == "found" - + call_args = mock_client.return_value.__aenter__.return_value.get.call_args assert call_args[1]["params"]["user.name"] == "john.doe" @@ -221,26 +148,26 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/entities", method="get", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func(**{"$filter": "name eq 'test'"}) assert result == "[]" - + call_args = mock_client.return_value.__aenter__.return_value.get.call_args assert call_args[1]["params"]["$filter"] == "name eq 'test'" @@ -257,26 +184,26 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/items", method="get", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func(**{"class": "premium"}) assert result == "items" - + call_args = mock_client.return_value.__aenter__.return_value.get.call_args assert call_args[1]["params"]["class"] == "premium" @@ -305,23 +232,23 @@ class TestCreateToolFunction: }, ] } - + func = create_tool_function( path="/repos/{repository-id}", method="get", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func( **{ "repository-id": "test-repo", @@ -347,26 +274,26 @@ class TestCreateToolFunction: }, } } - + func = create_tool_function( path="/create", method="post", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func(**{"body": {"name": "test"}}) assert result == "created" - + call_args = mock_client.return_value.__aenter__.return_value.post.call_args assert call_args[1]["json"] == {"name": "test"} @@ -374,23 +301,23 @@ class TestCreateToolFunction: async def test_no_parameters(self): """Test function with no parameters.""" operation = {} - + func = create_tool_function( path="/health", method="get", operation=operation, base_url="https://api.example.com", ) - + 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 ) - + result = await func() assert result == "ok" @@ -398,7 +325,7 @@ class TestCreateToolFunction: async def test_all_http_methods(self): """Test all supported HTTP methods.""" methods = ["get", "post", "put", "delete", "patch"] - + for method in methods: operation = { "parameters": [ @@ -410,20 +337,20 @@ class TestCreateToolFunction: } ] } - + func = create_tool_function( path="/repos/{repository-id}", method=method, operation=operation, base_url="https://api.example.com", ) - + 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 ) @@ -434,7 +361,7 @@ class TestCreateToolFunction: method, client_method, ) - + result = await func(**{"repository-id": "test"}) assert result == "success" @@ -442,20 +369,20 @@ class TestCreateToolFunction: """Verify that create_tool_function does not use exec().""" import ast import inspect - + # Get the source code of create_tool_function source = inspect.getsource(create_tool_function) - + # Parse the AST tree = ast.parse(source) - + # Check for exec() calls exec_calls = [] for node in ast.walk(tree): if isinstance(node, ast.Call): if isinstance(node.func, ast.Name) and node.func.id == "exec": exec_calls.append(node) - + # Should have no exec() calls assert len(exec_calls) == 0, "create_tool_function should not use exec()" @@ -487,14 +414,14 @@ class TestBuildInputSchema: }, ] } - + schema = build_input_schema(operation) - + # Original names should be in the schema assert "repository-id" in schema["properties"] assert "2fa-code" in schema["properties"] assert "$filter" in schema["properties"] - + # Required should include original names assert "repository-id" in schema["required"] @@ -514,9 +441,9 @@ class TestExtractParameters: "content": {"application/json": {"schema": {"type": "object"}}} }, } - + path_params, query_params, body_params = extract_parameters(operation) - + assert "repo-id" in path_params assert "filter" in query_params assert "data" in body_params @@ -525,4 +452,3 @@ class TestExtractParameters: if __name__ == "__main__": pytest.main([__file__, "-v"]) -