From 498c4254ead47876814b8a4f9b2fce485b943a37 Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sun, 7 Dec 2025 20:53:51 +0800 Subject: [PATCH 001/304] fix: Return 403 exception when calling GET responses api --- litellm/proxy/auth/auth_checks.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index fc79a4d3591..309bd577606 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -402,13 +402,14 @@ def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: - user_route: str - the route the user is trying to call - allowed_routes: List[str|LiteLLMRoutes] - the list of allowed routes for the user. """ + from starlette.routing import compile_path for allowed_route in allowed_routes: - if ( - allowed_route in LiteLLMRoutes.__members__ - and user_route in LiteLLMRoutes[allowed_route].value - ): - return True + if allowed_route in LiteLLMRoutes.__members__: + for template in LiteLLMRoutes[allowed_route].value: + regex, _, _ = compile_path(template) + if regex.match(user_route): + return True elif allowed_route == user_route: return True return False From 4ab58619ad56d93eb4add67bfd6064c50f449fa1 Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sun, 14 Dec 2025 17:55:13 +0800 Subject: [PATCH 002/304] fix: added new step into rotate master key function for processing credentials table --- .../proxy/credential_endpoints/endpoints.py | 9 ++--- .../key_management_endpoints.py | 34 +++++++++++++++++++ 2 files changed, 39 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 647abb73648..9f228bb1184 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -21,11 +21,11 @@ router = APIRouter() class CredentialHelperUtils: @staticmethod - def encrypt_credential_values(credential: CredentialItem) -> CredentialItem: + def encrypt_credential_values(credential: CredentialItem, new_encryption_key: Optional[str] = None) -> CredentialItem: """Encrypt values in credential.credential_values and add to DB""" encrypted_credential_values = {} for key, value in (credential.credential_values or {}).items(): - encrypted_credential_values[key] = encrypt_value_helper(value) + encrypted_credential_values[key] = encrypt_value_helper(value, new_encryption_key) # Return a new object to avoid mutating the caller's credential, which # is kept in memory and should remain unencrypted. @@ -246,7 +246,7 @@ async def delete_credential( def update_db_credential( - db_credential: CredentialItem, updated_patch: CredentialItem + db_credential: CredentialItem, updated_patch: CredentialItem, new_encryption_key: Optional[str] = None ) -> CredentialItem: """ Update a credential in the DB. @@ -258,7 +258,8 @@ def update_db_credential( ) encrypted_credential = CredentialHelperUtils.encrypt_credential_values( - updated_patch + updated_patch, + new_encryption_key, ) # update model name if encrypted_credential.credential_name: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index da44bda791d..8ea3122ce01 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2539,6 +2539,40 @@ async def _rotate_master_key( new_master_key=new_master_key, ) + # 5. process credentials table + try: + credentials = await prisma_client.db.litellm_credentialstable.find_many() + except Exception: + credentials = None + if credentials: + from litellm.proxy.credential_endpoints.endpoints import update_db_credential + + for cred in credentials: + try: + decrypted_cred = proxy_config.decrypt_credentials(cred) + encrypted_cred = update_db_credential( + db_credential=cred, + updated_patch=decrypted_cred, + new_encryption_key=new_master_key, + ) + credential_object_jsonified = jsonify_object(encrypted_cred.model_dump()) + await prisma_client.db.litellm_credentialstable.update( + where={"credential_name": cred.credential_name}, + data={ + **credential_object_jsonified, + "updated_by": user_api_key_dict.user_id, + }, + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to re-encrypt credential {cred.credential_name}: {str(e)}" + ) + # Continue with next credential instead of failing entire rotation + continue + verbose_proxy_logger.debug( + f"Successfully re-encrypted {len(credentials)} credentials with new master key" + ) + def get_new_token(data: Optional[RegenerateKeyRequest]) -> str: if data and data.new_key is not None: From f4f5ea85dfe7eff7130724b1d7331354468f5ce8 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 18 Dec 2025 14:42:41 +0530 Subject: [PATCH 003/304] Add redisvl in requirements.txt --- requirements.txt | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index c36f94f0752..2fb6c52cfd3 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,7 +7,8 @@ starlette==0.49.1 # starlette fastapi dep backoff==2.2.1 # server dep pyyaml==6.0.2 # server dep uvicorn==0.31.1 # server dep -gunicorn==23.0.0 # server dep +gunicorn==23.0.0 # server depredisvl +redisvl==0.4.1 # redis semantic cache fastuuid==0.13.5 # for uuid4 uvloop==0.21.0 # uvicorn dep, gives us much better performance under load boto3==1.36.0 # aws bedrock/sagemaker calls From 7ddab06bedba0d6151309c308a38cbfa13588dff Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sat, 20 Dec 2025 11:35:20 +0800 Subject: [PATCH 004/304] fix: fixed the issue of handling root paths when processing Discovery protected resource metadata and authorization server metadata URLs. --- .../mcp_server/discoverable_endpoints.py | 25 ++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ffa17a5b7c4..4b6020f582b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -15,6 +15,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.mcp_server.mcp_server_manager import MCPServer +from litellm.proxy.utils import get_server_root_path router = APIRouter( tags=["mcp"], @@ -381,7 +382,18 @@ async def callback(code: str, state: str): # ------------------------------ # Optional .well-known endpoints for MCP + OAuth discovery # ------------------------------ -@router.get("/.well-known/oauth-protected-resource/{mcp_server_name}/mcp") +""" + Per SEP-985, the client MUST: + 1. Try resource_metadata from WWW-Authenticate header (if present) + 2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path} + ( + If the resource identifier value contains a path or query component, any terminating slash (/) + following the host component MUST be removed before inserting /.well-known/ and the well-known + URI path suffix between the host component and the path(include root path) and/or query components. + https://datatracker.ietf.org/doc/html/rfc9728#section-3.1) + 3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource +""" +@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp") @router.get("/.well-known/oauth-protected-resource") async def oauth_protected_resource_mcp( request: Request, mcp_server_name: Optional[str] = None @@ -403,8 +415,15 @@ async def oauth_protected_resource_mcp( ), # this is what Claude will call } - -@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}") +""" + https://datatracker.ietf.org/doc/html/rfc8414#section-3.1 + RFC 8414: Path-aware OAuth discovery + If the issuer identifier value contains a path component, any + terminating "/" MUST be removed before inserting "/.well-known/" and + the well-known URI suffix between the host component and the path(include root path) + component. +""" +@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}") @router.get("/.well-known/oauth-authorization-server") async def oauth_authorization_server_mcp( request: Request, mcp_server_name: Optional[str] = None From 3a2ab6b0d12be8863d2a7604a1f3e8ab5721b521 Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sat, 20 Dec 2025 11:55:12 +0800 Subject: [PATCH 005/304] fix: added additional grant type into oauth_authorization_server response for fixing mcp auth register bad request issue --- .../proxy/_experimental/mcp_server/discoverable_endpoints.py | 2 +- .../_experimental/mcp_server/test_discoverable_endpoints.py | 3 ++- ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ffa17a5b7c4..5433196dfe3 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -428,7 +428,7 @@ async def oauth_authorization_server_mcp( "authorization_endpoint": authorization_endpoint, "token_endpoint": token_endpoint, "response_types_supported": ["code"], - "grant_types_supported": ["authorization_code"], + "grant_types_supported": ["authorization_code", "refresh_token"], "code_challenge_methods_supported": ["S256"], "token_endpoint_auth_methods_supported": ["client_secret_post"], # Claude expects a registration endpoint, even if we just fake it diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 6df9abd3fee..30f3d55f028 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -354,7 +354,7 @@ async def test_register_client_remote_registration_success(): request_payload = { "client_name": "Litellm Proxy", - "grant_types": ["authorization_code"], + "grant_types": ["authorization_code", "refresh_token"], "response_types": ["code"], "token_endpoint_auth_method": "client_secret_post", } @@ -603,6 +603,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): assert response["authorization_endpoint"].startswith("https://litellm.example.com/") assert response["token_endpoint"].startswith("https://litellm.example.com/") assert response["registration_endpoint"].startswith("https://litellm.example.com/") + assert response["grant_types_supported"] == ["authorization_code", "refresh_token"] @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx index 9600c962564..a62d8baa75b 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx @@ -136,7 +136,7 @@ export const useMcpOAuthFlow = ({ if (!hasPreconfiguredCredentials) { const registration = await registerMcpOAuthClient(accessToken, serverId, { client_name: temporaryPayload.alias || temporaryPayload.server_name || serverId, - grant_types: ["authorization_code"], + grant_types: ["authorization_code", "refresh_token"], response_types: ["code"], token_endpoint_auth_method: temporaryPayload.credentials && temporaryPayload.credentials.client_secret ? "client_secret_post" : "none", From 684fba42eaaf6a4d47795e56fd668b8d46b01525 Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sat, 20 Dec 2025 13:22:29 +0800 Subject: [PATCH 006/304] fix: added RFC RECOMMENDED property(scopes_supported) to protected resource and authorization server metadata --- .../mcp_server/discoverable_endpoints.py | 17 +++++- .../mcp_server/test_discoverable_endpoints.py | 54 ++++++++++++++++++- 2 files changed, 68 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d6fe3f2b9cf..ded591a8f53 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -398,8 +398,14 @@ async def callback(code: str, state: str): async def oauth_protected_resource_mcp( request: Request, mcp_server_name: Optional[str] = None ): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) # Get the correct base URL considering X-Forwarded-* headers request_base_url = get_request_base_url(request) + mcp_server: Optional[MCPServer] = None + if mcp_server_name: + mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name) return { "authorization_servers": [ ( @@ -413,6 +419,7 @@ async def oauth_protected_resource_mcp( if mcp_server_name else f"{request_base_url}/mcp" ), # this is what Claude will call + "scopes_supported": mcp_server.scopes if mcp_server else [], } """ @@ -428,6 +435,9 @@ async def oauth_protected_resource_mcp( async def oauth_authorization_server_mcp( request: Request, mcp_server_name: Optional[str] = None ): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) # Get the correct base URL considering X-Forwarded-* headers request_base_url = get_request_base_url(request) @@ -442,16 +452,21 @@ async def oauth_authorization_server_mcp( else f"{request_base_url}/token" ) + mcp_server: Optional[MCPServer] = None + if mcp_server_name: + mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name) + return { "issuer": request_base_url, # point to your proxy "authorization_endpoint": authorization_endpoint, "token_endpoint": token_endpoint, "response_types_supported": ["code"], + "scopes_supported": mcp_server.scopes if mcp_server else [], "grant_types_supported": ["authorization_code", "refresh_token"], "code_challenge_methods_supported": ["S256"], "token_endpoint_auth_methods_supported": ["client_secret_post"], # Claude expects a registration endpoint, even if we just fake it - "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register", + "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register" if mcp_server_name else f"{request_base_url}/register", } diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 30f3d55f028..4c5723b8284 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -556,9 +556,33 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( oauth_protected_resource_mcp, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) @@ -568,13 +592,14 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): # Call the endpoint response = await oauth_protected_resource_mcp( request=mock_request, - mcp_server_name="test_server", + mcp_server_name="test_oauth", ) # Verify response uses HTTPS URLs assert response["authorization_servers"][0].startswith( "https://litellm.example.com/" ) + assert response["scopes_supported"] == oauth2_server.scopes @pytest.mark.asyncio @@ -584,9 +609,33 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( oauth_authorization_server_mcp, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) @@ -596,7 +645,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): # Call the endpoint response = await oauth_authorization_server_mcp( request=mock_request, - mcp_server_name="test_server", + mcp_server_name="test_oauth", ) # Verify response uses HTTPS URLs @@ -604,6 +653,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): assert response["token_endpoint"].startswith("https://litellm.example.com/") assert response["registration_endpoint"].startswith("https://litellm.example.com/") assert response["grant_types_supported"] == ["authorization_code", "refresh_token"] + assert response["scopes_supported"] == oauth2_server.scopes @pytest.mark.asyncio From 0306f02e74d7fff462fb727639a60e5ae11d64e4 Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sat, 20 Dec 2025 14:13:18 +0800 Subject: [PATCH 007/304] fix: removed initialize the tool name to MCP server name mapping(oauth2) on startup for avoiding 401 error --- .../mcp_server/mcp_server_manager.py | 3 +++ .../mcp_server/test_mcp_server_manager.py | 19 +++++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8c9d8630457..c2215efe9d0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1913,6 +1913,9 @@ class MCPServerManager: Note: This now handles prefixed tool names """ for server in self.get_registry().values(): + if server.auth_type == MCPAuth.oauth2: + # Skip OAuth2 servers for now as they may require user-specific tokens + continue tools = await self._get_tools_from_server(server) for tool in tools: # The tool.name here is already prefixed from _get_tools_from_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 7a6e5ad17f6..c0ded9c728c 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 @@ -536,7 +536,26 @@ class TestMCPServerManager: assert ( server.registration_url == "https://discovered.example.com/register" ) + @pytest.mark.asyncio + async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self): + manager = MCPServerManager() + config = { + "example": { + "url": "https://example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "scopes": ["config"], + "authorization_url": "https://config.example.com/auth", + } + } + + await manager.load_servers_from_config(config) + + # Initialize the tool mapping + await manager._initialize_tool_name_to_mcp_server_name_mapping() + assert manager.tool_name_to_mcp_server_name_mapping == {} + @pytest.mark.asyncio async def test_list_tools_handles_missing_server_alias(self): """Test that list_tools handles servers without alias gracefully""" From 36a369a747fccedb87e12cf26996aebfc221f57e Mon Sep 17 00:00:00 2001 From: Eric84626 Date: Sat, 20 Dec 2025 14:28:52 +0800 Subject: [PATCH 008/304] fix: upgraded mcp sdk depency version for fixing ClosedResourceError --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index f222acc46e6..cb12a658814 100644 --- a/requirements.txt +++ b/requirements.txt @@ -20,7 +20,7 @@ google-cloud-aiplatform==1.47.0 # for vertex ai calls google-cloud-iam==2.19.1 # for GCP IAM Redis authentication google-genai==1.22.0 anthropic[vertex]==0.54.0 -mcp==1.21.2 ; python_version >= "3.10" # for MCP server +mcp==1.25.0 ; python_version >= "3.10" # for MCP server google-generativeai==0.5.0 # for vertex ai calls async_generator==1.10.0 # for async ollama calls langfuse==2.59.7 # for langfuse self-hosted logging From b33f1ec2b215bd97d800a98004e6aeaa5d8d21ce Mon Sep 17 00:00:00 2001 From: hamzaq453 Date: Sun, 28 Dec 2025 19:28:11 +0500 Subject: [PATCH 009/304] Fix: Remove exec() usage and handle invalid OpenAPI parameter names - Add to_safe_identifier() to convert any parameter name to valid Python identifier - Refactor create_tool_function() to use closure with **kwargs instead of exec() - Handle edge cases: hyphens, dots, leading digits, Python keywords, special chars - Add comprehensive test suite covering all edge cases - Fixes #18471: OpenAPI MCP server crashes on invalid parameter names - Security: Eliminates arbitrary code execution risk from untrusted OpenAPI specs --- .../mcp_server/openapi_to_mcp_generator.py | 234 +++++--- .../test_openapi_to_mcp_generator.py | 528 ++++++++++++++++++ 2 files changed, 695 insertions(+), 67 deletions(-) create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py 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 72288f8e673..dc5d0ce73af 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -3,6 +3,8 @@ 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 @@ -17,6 +19,64 @@ 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: @@ -112,12 +172,20 @@ 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. + Args: path: API endpoint path method: HTTP method (get, post, put, delete, patch) operation: OpenAPI operation object base_url: Base URL for the API headers: Optional headers to include in requests (e.g., authentication) + + Returns: + An async function that accepts **kwargs and makes the HTTP request """ if headers is None: headers = {} @@ -125,77 +193,109 @@ def create_tool_function( path_params, query_params, body_params = extract_parameters(operation) all_params = path_params + query_params + body_params - # Build function signature dynamically - if all_params: - params_str = ", ".join(f"{p}: str = ''" for p in all_params) - else: - params_str = "" - - # Create the function code as a string - func_code = f''' -async def tool_function({params_str}) -> str: - """Dynamically generated tool function.""" - url = base_url + path + # 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] = {} - # Replace path parameters - path_param_names = {path_params} - for param_name in path_param_names: - param_value = locals().get(param_name, "") - if param_value: - url = url.replace("{{" + param_name + "}}", str(param_value)) - - # Build query params - query_param_names = {query_params} - params = {{}} - for param_name in query_param_names: - param_value = locals().get(param_name, "") - if param_value: - params[param_name] = param_value - - # Build request body - body_param_names = {body_params} - json_body = None - if body_param_names: - body_value = locals().get("body", {{}}) - if isinstance(body_value, dict): - json_body = body_value - elif body_value: - # If it's a string, try to parse as JSON - import json as json_module - try: - json_body = json_module.loads(body_value) if isinstance(body_value, str) else {{"data": body_value}} - except: - json_body = {{"data": body_value}} - - # Make HTTP request - async with httpx.AsyncClient() as client: - 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 "Unsupported HTTP method: {method}" + 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 - return response.text -''' + param_name_map[orig_name] = safe_name + safe_to_original_map[safe_name] = orig_name - # Execute the function code to create the actual function - local_vars = { - "httpx": httpx, - "headers": headers, - "base_url": base_url, - "path": path, - "method": method, - } - exec(func_code, local_vars) + # 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() - return local_vars["tool_function"] + 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. + """ + # 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, "") + + 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)) + + # 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, "") + + if param_value: + # Use original parameter name in query string (as expected by API) + params[orig_param_name] = param_value + + # Build request body + json_body: Optional[Dict[str, Any]] = None + if original_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, {}) + 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} + 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}" + + return response.text + + return tool_function def register_tools_from_openapi(spec: Dict[str, Any], base_url: str): 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 new file mode 100644 index 00000000000..fa28dc34981 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -0,0 +1,528 @@ +""" +Tests for OpenAPI to MCP generator, focusing on security and edge cases. + +This test suite ensures that: +1. Parameter names with invalid Python identifiers are handled safely +2. No exec() is used (security) +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.""" + + @pytest.mark.asyncio + async def test_hyphenated_path_parameter(self): + """Test function with hyphenated path parameter (e.g., repository-id).""" + operation = { + "parameters": [ + { + "name": "repository-id", + "in": "path", + "required": True, + "schema": {"type": "string"}, + } + ] + } + + 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() + mock_response.text = '{"id": "123"}' + 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]) + + @pytest.mark.asyncio + async def test_leading_digit_parameter(self): + """Test function with parameter starting with digit (e.g., 2fa-code).""" + operation = { + "parameters": [ + { + "name": "2fa-code", + "in": "query", + "required": False, + "schema": {"type": "string"}, + } + ] + } + + 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" + + @pytest.mark.asyncio + async def test_dot_in_parameter_name(self): + """Test function with dot in parameter name (e.g., user.name).""" + operation = { + "parameters": [ + { + "name": "user.name", + "in": "query", + "required": False, + "schema": {"type": "string"}, + } + ] + } + + 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" + + @pytest.mark.asyncio + async def test_dollar_sign_parameter(self): + """Test function with dollar sign parameter (OData style, e.g., $filter).""" + operation = { + "parameters": [ + { + "name": "$filter", + "in": "query", + "required": False, + "schema": {"type": "string"}, + } + ] + } + + 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'" + + @pytest.mark.asyncio + async def test_python_keyword_parameter(self): + """Test function with Python keyword as parameter name (e.g., class).""" + operation = { + "parameters": [ + { + "name": "class", + "in": "query", + "required": False, + "schema": {"type": "string"}, + } + ] + } + + 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" + + @pytest.mark.asyncio + async def test_multiple_problematic_parameters(self): + """Test function with multiple problematic parameter names.""" + operation = { + "parameters": [ + { + "name": "repository-id", + "in": "path", + "required": True, + "schema": {"type": "string"}, + }, + { + "name": "2fa-code", + "in": "query", + "required": False, + "schema": {"type": "string"}, + }, + { + "name": "$filter", + "in": "query", + "required": False, + "schema": {"type": "string"}, + }, + ] + } + + 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", + "2fa-code": "123", + "$filter": "active", + } + ) + assert result == "success" + + @pytest.mark.asyncio + async def test_request_body_parameter(self): + """Test function with request body parameter.""" + operation = { + "requestBody": { + "required": True, + "content": { + "application/json": { + "schema": { + "type": "object", + "properties": {"name": {"type": "string"}}, + } + } + }, + } + } + + 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"} + + @pytest.mark.asyncio + 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" + + @pytest.mark.asyncio + async def test_all_http_methods(self): + """Test all supported HTTP methods.""" + methods = ["get", "post", "put", "delete", "patch"] + + for method in methods: + operation = { + "parameters": [ + { + "name": "repository-id", + "in": "path", + "required": True, + "schema": {"type": "string"}, + } + ] + } + + 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 + ) + client_method.return_value = mock_response + client_method = AsyncMock(return_value=mock_response) + setattr( + mock_client.return_value.__aenter__.return_value, + method, + client_method, + ) + + result = await func(**{"repository-id": "test"}) + assert result == "success" + + def test_no_exec_usage(self): + """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()" + + +class TestBuildInputSchema: + """Test that build_input_schema preserves original parameter names.""" + + def test_original_parameter_names_preserved(self): + """Test that original parameter names are preserved in input schema.""" + operation = { + "parameters": [ + { + "name": "repository-id", + "in": "path", + "required": True, + "schema": {"type": "string"}, + }, + { + "name": "2fa-code", + "in": "query", + "required": False, + "schema": {"type": "string"}, + }, + { + "name": "$filter", + "in": "query", + "required": False, + "schema": {"type": "string"}, + }, + ] + } + + 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"] + + +class TestExtractParameters: + """Test parameter extraction from OpenAPI operations.""" + + def test_extract_path_query_body_params(self): + """Test extraction of different parameter types.""" + operation = { + "parameters": [ + {"name": "repo-id", "in": "path"}, + {"name": "filter", "in": "query"}, + {"name": "data", "in": "body"}, + ], + "requestBody": { + "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 + assert "body" in body_params # From requestBody + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) + From 4573ab326b5b4293261ab7795b056018c06d1fa7 Mon Sep 17 00:00:00 2001 From: hamzaq453 Date: Tue, 30 Dec 2025 14:04:15 +0500 Subject: [PATCH 010/304] 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"]) - From 78693bb9d065924c2b4fd26962aeedc8eb84f208 Mon Sep 17 00:00:00 2001 From: mangabits <1457532+mangabits@users.noreply.github.com> Date: Fri, 19 Dec 2025 18:06:02 -0800 Subject: [PATCH 011/304] Use already configured opentelemetry providers Users that instrument using opentelemetry-instrument can now setup exporters as per their environment. --- litellm/integrations/opentelemetry.py | 152 ++++++++++++------ .../test_opentelemetry_unit_tests.py | 25 --- .../integrations/test_opentelemetry.py | 92 +++++++++++ 3 files changed, 192 insertions(+), 77 deletions(-) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 12e60bc25bb..c874c799d5e 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -196,50 +196,87 @@ class OpenTelemetry(CustomLogger): litellm.service_callback.append(self) setattr(proxy_server, "open_telemetry_logger", self) + def _get_or_create_provider( + self, + provider, + provider_name: str, + get_existing_provider_fn, + sdk_provider_class, + create_new_provider_fn, + set_provider_fn, + ): + """ + Generic helper to get or create an OpenTelemetry provider (Tracer, Meter, or Logger). + + Args: + provider: The provider instance passed to the init function (can be None) + provider_name: Name for logging (e.g., "TracerProvider") + get_existing_provider_fn: Function to get the existing global provider + sdk_provider_class: The SDK provider class to check for (e.g., TracerProvider from SDK) + create_new_provider_fn: Function to create a new provider instance + set_provider_fn: Function to set the provider globally + + Returns: + The provider to use (either existing, new, or explicitly provided) + """ + if provider is None: + # Check if a provider is already set globally + try: + existing_provider = get_existing_provider_fn() + + # If a real SDK provider exists (set by another SDK like Langfuse), use it + # This uses a positive check for SDK providers instead of a negative check for proxy providers + if isinstance(existing_provider, sdk_provider_class): + verbose_logger.debug( + "OpenTelemetry: Using existing %s: %s", + provider_name, + type(existing_provider).__name__, + ) + provider = existing_provider + # Don't call set_provider to preserve existing context + else: + # Default proxy provider or unknown type, create our own + verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name) + provider = create_new_provider_fn() + set_provider_fn(provider) + except Exception as e: + # Fallback: create a new provider if something goes wrong + verbose_logger.debug( + "OpenTelemetry: Exception checking existing %s, creating new one: %s", + provider_name, + str(e), + ) + provider = create_new_provider_fn() + set_provider_fn(provider) + else: + # Provider explicitly provided (e.g., for testing) + # Do NOT call set_provider_fn - the caller is responsible for managing global state + # If they want it to be global, they've already set it before passing it to us + verbose_logger.debug( + "OpenTelemetry: Using provided TracerProvider: %s", + type(provider).__name__, + ) + + return provider + def _init_tracing(self, tracer_provider): from opentelemetry import trace from opentelemetry.sdk.trace import TracerProvider from opentelemetry.trace import SpanKind - # use provided tracer or create a new one - if tracer_provider is None: - # Check if a TracerProvider is already set globally (e.g., by Langfuse SDK) - try: - from opentelemetry.trace import ProxyTracerProvider + def create_tracer_provider(): + provider = TracerProvider(resource=_get_litellm_resource()) + provider.add_span_processor(self._get_span_processor()) + return provider - existing_provider = trace.get_tracer_provider() - - # If an actual provider exists (not the default proxy), use it - if not isinstance(existing_provider, ProxyTracerProvider): - verbose_logger.debug( - "OpenTelemetry: Using existing TracerProvider: %s", - type(existing_provider).__name__, - ) - tracer_provider = existing_provider - # Don't call set_tracer_provider to preserve existing context - else: - # No real provider exists yet, create our own - verbose_logger.debug("OpenTelemetry: Creating new TracerProvider") - tracer_provider = TracerProvider(resource=_get_litellm_resource()) - tracer_provider.add_span_processor(self._get_span_processor()) - trace.set_tracer_provider(tracer_provider) - except Exception as e: - # Fallback: create a new provider if something goes wrong - verbose_logger.debug( - "OpenTelemetry: Exception checking existing provider, creating new one: %s", - str(e), - ) - tracer_provider = TracerProvider(resource=_get_litellm_resource()) - tracer_provider.add_span_processor(self._get_span_processor()) - trace.set_tracer_provider(tracer_provider) - else: - # Tracer provider explicitly provided (e.g., for testing) - # Do NOT call set_tracer_provider - the caller is responsible for managing global state - # If they want it to be global, they've already set it before passing it to us - verbose_logger.debug( - "OpenTelemetry: Using provided TracerProvider: %s", - type(tracer_provider).__name__, - ) + tracer_provider = self._get_or_create_provider( + provider=tracer_provider, + provider_name="TracerProvider", + get_existing_provider_fn=trace.get_tracer_provider, + sdk_provider_class=TracerProvider, + create_new_provider_fn=create_tracer_provider, + set_provider_fn=trace.set_tracer_provider, + ) # Grab our tracer from the TracerProvider (not from global context) # This ensures we use the provided TracerProvider (e.g., for testing) @@ -259,8 +296,7 @@ class OpenTelemetry(CustomLogger): from opentelemetry import metrics from opentelemetry.sdk.metrics import Histogram, MeterProvider - # Only create OTLP infrastructure if no custom meter provider is provided - if meter_provider is None: + def create_meter_provider(): from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( OTLPMetricExporter, ) @@ -281,15 +317,20 @@ class OpenTelemetry(CustomLogger): _metric_exporter, export_interval_millis=10000 ) - meter_provider = MeterProvider( + return MeterProvider( metric_readers=[_metric_reader], resource=_get_litellm_resource() ) - meter = meter_provider.get_meter(__name__) - else: - # Use the provided meter provider as-is, without creating additional OTLP infrastructure - meter = meter_provider.get_meter(__name__) - metrics.set_meter_provider(meter_provider) + meter_provider = self._get_or_create_provider( + provider=meter_provider, + provider_name="MeterProvider", + get_existing_provider_fn=metrics.get_meter_provider, + sdk_provider_class=MeterProvider, + create_new_provider_fn=create_meter_provider, + set_provider_fn=metrics.set_meter_provider, + ) + + meter = meter_provider.get_meter(__name__) self._operation_duration_histogram = meter.create_histogram( name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38 @@ -327,22 +368,29 @@ class OpenTelemetry(CustomLogger): if not self.config.enable_events: return - from opentelemetry._logs import set_logger_provider + from opentelemetry._logs import get_logger_provider, set_logger_provider from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider from opentelemetry.sdk._logs.export import BatchLogRecordProcessor - # set up log pipeline - if logger_provider is None: + def create_logger_provider(): litellm_resource = _get_litellm_resource() - logger_provider = OTLoggerProvider(resource=litellm_resource) + provider = OTLoggerProvider(resource=litellm_resource) # Only add OTLP exporter if we created the logger provider ourselves log_exporter = self._get_log_exporter() if log_exporter: - logger_provider.add_log_record_processor( + provider.add_log_record_processor( BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type] ) + return provider - set_logger_provider(logger_provider) + logger_provider = self._get_or_create_provider( + provider=logger_provider, + provider_name="LoggerProvider", + get_existing_provider_fn=get_logger_provider, + sdk_provider_class=OTLoggerProvider, + create_new_provider_fn=create_logger_provider, + set_provider_fn=set_logger_provider, + ) def log_success_event(self, kwargs, response_obj, start_time, end_time): self._handle_success(kwargs, response_obj, start_time, end_time) diff --git a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py index 3d0682d9033..04f8abe64de 100644 --- a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py +++ b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py @@ -65,31 +65,6 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest): # External spans should only be closed by their creators parent_otel_span.end.assert_not_called() - def test_init_tracing_respects_existing_tracer_provider(self): - """ - Unit test: _init_tracing() should respect existing TracerProvider. - - When a TracerProvider already exists (e.g., set by Langfuse SDK), - LiteLLM should use it instead of creating a new one. - """ - from opentelemetry import trace - from opentelemetry.sdk.trace import TracerProvider - from litellm.integrations.opentelemetry import OpenTelemetry - - # Setup: Create and set an existing TracerProvider - tracer_provider = TracerProvider() - trace.set_tracer_provider(tracer_provider) - existing_provider = trace.get_tracer_provider() - - # Act: Initialize OpenTelemetry integration (should detect existing provider) - otel_integration = OpenTelemetry() - - # Assert: The existing provider should still be active - current_provider = trace.get_tracer_provider() - assert current_provider is existing_provider, ( - "Existing TracerProvider should be respected and not overridden" - ) - def test_get_span_context_detects_active_span(self): """ Unit test: _get_span_context() should auto-detect active spans from global context. diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 5d648f601f6..bd7d4bf2179 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -172,6 +172,98 @@ class TestOpenTelemetryCostBreakdown(unittest.TestCase): assert ("gen_ai.cost.original_cost", 0.004) not in call_args_list +class TestOpenTelemetryProviderInitialization(unittest.TestCase): + """Test suite for verifying provider initialization respects existing providers""" + + def test_init_tracing_respects_existing_tracer_provider(self): + """ + Unit test: _init_tracing() should respect existing TracerProvider. + + When a TracerProvider already exists (e.g., set by Langfuse SDK), + LiteLLM should use it instead of creating a new one. + """ + from opentelemetry import trace + from opentelemetry.sdk.trace import TracerProvider + + # Setup: Create and set an existing TracerProvider + tracer_provider = TracerProvider() + trace.set_tracer_provider(tracer_provider) + existing_provider = trace.get_tracer_provider() + + # Act: Initialize OpenTelemetry integration (should detect existing provider) + otel_integration = OpenTelemetry() + + # Assert: The existing provider should still be active + current_provider = trace.get_tracer_provider() + assert current_provider is existing_provider, ( + "Existing TracerProvider should be respected and not overridden" + ) + + def test_init_metrics_respects_existing_meter_provider(self): + """ + Unit test: _init_metrics() should respect existing MeterProvider. + + When a MeterProvider already exists (e.g., set by Langfuse SDK), + LiteLLM should use it instead of creating a new one. + """ + from opentelemetry import metrics + from opentelemetry.sdk.metrics import MeterProvider + + # Setup: Enable metrics for this test + os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true" + + try: + # Create and set an existing MeterProvider + meter_provider = MeterProvider() + metrics.set_meter_provider(meter_provider) + existing_provider = metrics.get_meter_provider() + + # Act: Initialize OpenTelemetry integration (should detect existing provider) + config = OpenTelemetryConfig.from_env() + otel_integration = OpenTelemetry(config=config) + + # Assert: The existing provider should still be active + current_provider = metrics.get_meter_provider() + assert current_provider is existing_provider, ( + "Existing MeterProvider should be respected and not overridden" + ) + finally: + # Cleanup + os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_METRICS", None) + + def test_init_logs_respects_existing_logger_provider(self): + """ + Unit test: _init_logs() should respect existing LoggerProvider. + + When a LoggerProvider already exists (e.g., set by Langfuse SDK), + LiteLLM should use it instead of creating a new one. + """ + from opentelemetry._logs import get_logger_provider, set_logger_provider + from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider + + # Setup: Enable events for this test + os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS"] = "true" + + try: + # Create and set an existing LoggerProvider + logger_provider = OTLoggerProvider() + set_logger_provider(logger_provider) + existing_provider = get_logger_provider() + + # Act: Initialize OpenTelemetry integration (should detect existing provider) + config = OpenTelemetryConfig.from_env() + otel_integration = OpenTelemetry(config=config) + + # Assert: The existing provider should still be active + current_provider = get_logger_provider() + assert current_provider is existing_provider, ( + "Existing LoggerProvider should be respected and not overridden" + ) + finally: + # Cleanup + os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) + + class TestOpenTelemetry(unittest.TestCase): POLL_INTERVAL = 0.05 POLL_TIMEOUT = 2.0 From 9ca043f1d85020e779bb2d907683ce816534d02c Mon Sep 17 00:00:00 2001 From: mangabits <1457532+mangabits@users.noreply.github.com> Date: Mon, 29 Dec 2025 17:21:28 -0800 Subject: [PATCH 012/304] Handle all protocols for all telemetry --- litellm/integrations/opentelemetry.py | 107 ++++++++++++++++++++------ 1 file changed, 82 insertions(+), 25 deletions(-) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index c874c799d5e..b356e41c152 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -294,32 +294,18 @@ class OpenTelemetry(CustomLogger): return from opentelemetry import metrics - from opentelemetry.sdk.metrics import Histogram, MeterProvider + from opentelemetry.sdk.metrics import MeterProvider def create_meter_provider(): - from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( - OTLPMetricExporter, - ) - from opentelemetry.sdk.metrics.export import ( - AggregationTemporality, - PeriodicExportingMetricReader, - ) - - normalized_endpoint = self._normalize_otel_endpoint( - self.config.endpoint, "metrics" - ) - _metric_exporter = OTLPMetricExporter( - endpoint=normalized_endpoint, - headers=OpenTelemetry._get_headers_dictionary(self.config.headers), - preferred_temporality={Histogram: AggregationTemporality.DELTA}, - ) - _metric_reader = PeriodicExportingMetricReader( - _metric_exporter, export_interval_millis=10000 - ) - - return MeterProvider( - metric_readers=[_metric_reader], resource=_get_litellm_resource() + metric_reader = self._get_metric_reader() + if metric_reader: + return MeterProvider( + metric_readers=[metric_reader], resource=_get_litellm_resource() + ) + verbose_logger.warning( + "OpenTelemetry: No metric reader created. Metrics will not be exported." ) + return MeterProvider(resource=_get_litellm_resource()) meter_provider = self._get_or_create_provider( provider=meter_provider, @@ -383,7 +369,7 @@ class OpenTelemetry(CustomLogger): ) return provider - logger_provider = self._get_or_create_provider( + self._get_or_create_provider( provider=logger_provider, provider_name="LoggerProvider", get_existing_provider_fn=get_logger_provider, @@ -992,6 +978,15 @@ class OpenTelemetry(CustomLogger): if not self.config.enable_events: return + # NOTE: Semantic logs (gen_ai.content.prompt/completion events) have compatibility issues + # with OTEL SDK >= 1.39.0 due to breaking changes in PR #4676: + # - LogRecord moved from opentelemetry.sdk._logs to opentelemetry.sdk._logs._internal + # - LogRecord constructor no longer accepts 'resource' parameter (now inherited from LoggerProvider) + # - LogData class was removed entirely + # These logs work correctly in OTEL SDK < 1.39.0 but may fail in >= 1.39.0. + # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 + # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords + from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider from opentelemetry.sdk._logs import LogRecord as SdkLogRecord @@ -1855,7 +1850,8 @@ class OpenTelemetry(CustomLogger): ) return self.OTEL_EXPORTER - if self.OTEL_EXPORTER == "console": + otel_logs_exporter = os.getenv("OTEL_LOGS_EXPORTER") + if self.OTEL_EXPORTER == "console" or otel_logs_exporter == "console": from opentelemetry.sdk._logs.export import ConsoleLogExporter verbose_logger.debug( @@ -1902,6 +1898,67 @@ class OpenTelemetry(CustomLogger): return ConsoleLogExporter() + def _get_metric_reader(self): + """ + Get the appropriate metric reader based on the configuration. + """ + from opentelemetry.sdk.metrics import Histogram + from opentelemetry.sdk.metrics.export import ( + AggregationTemporality, + ConsoleMetricExporter, + PeriodicExportingMetricReader, + ) + + verbose_logger.debug( + "OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s", + self.OTEL_EXPORTER, + self.OTEL_ENDPOINT, + self.OTEL_HEADERS, + ) + + _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS) + normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "metrics") + + if self.OTEL_EXPORTER == "console": + exporter = ConsoleMetricExporter() + return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + + elif ( + self.OTEL_EXPORTER == "otlp_http" + or self.OTEL_EXPORTER == "http/protobuf" + or self.OTEL_EXPORTER == "http/json" + ): + from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( + OTLPMetricExporter, + ) + + exporter = OTLPMetricExporter( + endpoint=normalized_endpoint, + headers=_split_otel_headers, + preferred_temporality={Histogram: AggregationTemporality.DELTA}, + ) + return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + + elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc": + from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( + OTLPMetricExporter, + ) + + exporter = OTLPMetricExporter( + endpoint=normalized_endpoint, + headers=_split_otel_headers, + preferred_temporality={Histogram: AggregationTemporality.DELTA}, + ) + return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + + else: + verbose_logger.warning( + "OpenTelemetry: Unknown metric exporter '%s', defaulting to console. Supported: console, otlp_http, otlp_grpc", + self.OTEL_EXPORTER, + ) + exporter = ConsoleMetricExporter() + return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + def _normalize_otel_endpoint( self, endpoint: Optional[str], signal_type: str ) -> Optional[str]: From 4a1bc90e61a47512e7355ace5f246e89672f2cdd Mon Sep 17 00:00:00 2001 From: mangabits <1457532+mangabits@users.noreply.github.com> Date: Tue, 30 Dec 2025 18:11:54 -0800 Subject: [PATCH 013/304] Add more tests --- .../integrations/test_opentelemetry.py | 87 ++++++++++--------- 1 file changed, 44 insertions(+), 43 deletions(-) diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index bd7d4bf2179..6c17570e135 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -199,6 +199,7 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): "Existing TracerProvider should be respected and not overridden" ) + @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True) def test_init_metrics_respects_existing_meter_provider(self): """ Unit test: _init_metrics() should respect existing MeterProvider. @@ -209,28 +210,22 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): from opentelemetry import metrics from opentelemetry.sdk.metrics import MeterProvider - # Setup: Enable metrics for this test - os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true" + # Create and set an existing MeterProvider + meter_provider = MeterProvider() + metrics.set_meter_provider(meter_provider) + existing_provider = metrics.get_meter_provider() - try: - # Create and set an existing MeterProvider - meter_provider = MeterProvider() - metrics.set_meter_provider(meter_provider) - existing_provider = metrics.get_meter_provider() + # Act: Initialize OpenTelemetry integration (should detect existing provider) + config = OpenTelemetryConfig.from_env() + otel_integration = OpenTelemetry(config=config) - # Act: Initialize OpenTelemetry integration (should detect existing provider) - config = OpenTelemetryConfig.from_env() - otel_integration = OpenTelemetry(config=config) - - # Assert: The existing provider should still be active - current_provider = metrics.get_meter_provider() - assert current_provider is existing_provider, ( - "Existing MeterProvider should be respected and not overridden" - ) - finally: - # Cleanup - os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_METRICS", None) + # Assert: The existing provider should still be active + current_provider = metrics.get_meter_provider() + assert current_provider is existing_provider, ( + "Existing MeterProvider should be respected and not overridden" + ) + @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS": "true"}, clear=True) def test_init_logs_respects_existing_logger_provider(self): """ Unit test: _init_logs() should respect existing LoggerProvider. @@ -241,27 +236,20 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): from opentelemetry._logs import get_logger_provider, set_logger_provider from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider - # Setup: Enable events for this test - os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS"] = "true" + # Create and set an existing LoggerProvider + logger_provider = OTLoggerProvider() + set_logger_provider(logger_provider) + existing_provider = get_logger_provider() - try: - # Create and set an existing LoggerProvider - logger_provider = OTLoggerProvider() - set_logger_provider(logger_provider) - existing_provider = get_logger_provider() + # Act: Initialize OpenTelemetry integration (should detect existing provider) + config = OpenTelemetryConfig.from_env() + otel_integration = OpenTelemetry(config=config) - # Act: Initialize OpenTelemetry integration (should detect existing provider) - config = OpenTelemetryConfig.from_env() - otel_integration = OpenTelemetry(config=config) - - # Assert: The existing provider should still be active - current_provider = get_logger_provider() - assert current_provider is existing_provider, ( - "Existing LoggerProvider should be respected and not overridden" - ) - finally: - # Cleanup - os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) + # Assert: The existing provider should still be active + current_provider = get_logger_provider() + assert current_provider is existing_provider, ( + "Existing LoggerProvider should be respected and not overridden" + ) class TestOpenTelemetry(unittest.TestCase): @@ -712,7 +700,6 @@ class TestOpenTelemetry(unittest.TestCase): self.assertEqual(attributes.get("extra.attr"), "extra-value") - def test_handle_success_spans_only(self): # make sure neither events nor metrics is on os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) @@ -779,11 +766,8 @@ class TestOpenTelemetry(unittest.TestCase): logs = log_exporter.get_finished_logs() self.assertFalse(logs, "Did not expect any logs") + @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True) def test_handle_success_spans_and_metrics(self): - # only metrics on - os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) - os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true" - # ─── build in‐memory OTEL providers/exporters ───────────────────────────── span_exporter = InMemorySpanExporter() tracer_provider = TracerProvider() @@ -1412,6 +1396,23 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): ) self.assertEqual(normalized, "http://collector:4317/v1/logs") + def test_get_metric_reader_uses_http_exporter_for_http_protobuf(self): + """Test that http/protobuf protocol uses OTLPMetricExporterHTTP""" + from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( + OTLPMetricExporter, + ) + from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader + + config = OpenTelemetryConfig( + exporter="http/protobuf", endpoint="http://collector:4318" + ) + otel = OpenTelemetry(config=config) + + reader = otel._get_metric_reader() + + self.assertIsInstance(reader, PeriodicExportingMetricReader) + self.assertIsInstance(reader._exporter, OTLPMetricExporter) + class TestOpenTelemetryExternalSpan(unittest.TestCase): """ From da05b56756d0c8144b9c20213c0016b9a3a8f157 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Fri, 2 Jan 2026 17:07:52 +0900 Subject: [PATCH 014/304] feat: add user email to cloudzero --- litellm/integrations/cloudzero/cloudzero.py | 3 ++ litellm/integrations/cloudzero/database.py | 4 +- litellm/integrations/cloudzero/transform.py | 3 +- .../cloudzero/test_dry_run_endpoint.py | 6 ++- .../integrations/cloudzero/test_transform.py | 52 ++++++++++++++++++- 5 files changed, 63 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/cloudzero/cloudzero.py b/litellm/integrations/cloudzero/cloudzero.py index 403829deba0..9da8ea52b5c 100644 --- a/litellm/integrations/cloudzero/cloudzero.py +++ b/litellm/integrations/cloudzero/cloudzero.py @@ -317,6 +317,7 @@ class CloudZeroLogger(CustomLogger): ) cbf_table.add_column("team_id", style="cyan", no_wrap=False) cbf_table.add_column("team_alias", style="cyan", no_wrap=False) + cbf_table.add_column("user_email", style="cyan", no_wrap=False) cbf_table.add_column("api_key_alias", style="yellow", no_wrap=False) cbf_table.add_column( "usage/amount", style="yellow", justify="right", no_wrap=False @@ -339,6 +340,7 @@ class CloudZeroLogger(CustomLogger): entity_id = str(record.get("entity_id", "N/A")) team_id = str(record.get("resource/tag:team_id", "N/A")) team_alias = str(record.get("resource/tag:team_alias", "N/A")) + user_email = str(record.get("resource/tag:user_email", "N/A")) api_key_alias = str(record.get("resource/tag:api_key_alias", "N/A")) cbf_table.add_row( @@ -348,6 +350,7 @@ class CloudZeroLogger(CustomLogger): entity_id, team_id, team_alias, + user_email, api_key_alias, usage_amount, resource_id, diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 83ca01a5c0e..2128b55bf83 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -79,10 +79,12 @@ class LiteLLMDatabase: dus.updated_at, vt.team_id, vt.key_alias as api_key_alias, - tt.team_alias + tt.team_alias, + ut.user_email as user_email FROM "LiteLLM_DailyUserSpend" dus LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id + LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id {where_clause} ORDER BY dus.date DESC, dus.created_at DESC """ diff --git a/litellm/integrations/cloudzero/transform.py b/litellm/integrations/cloudzero/transform.py index e0263295388..e06b944a419 100644 --- a/litellm/integrations/cloudzero/transform.py +++ b/litellm/integrations/cloudzero/transform.py @@ -98,6 +98,7 @@ class CBFTransformer: # Handle team information with fallbacks team_id = row.get('team_id') team_alias = row.get('team_alias') + user_email = row.get('user_email') # Use team_alias if available, otherwise team_id, otherwise fallback to 'unknown' entity_id = str(team_alias) if team_alias else (str(team_id) if team_id else 'unknown') @@ -112,6 +113,7 @@ class CBFTransformer: 'provider': str(row.get('custom_llm_provider', '')), 'api_key_prefix': api_key_hash, 'api_key_alias': str(row.get('api_key_alias', '')), + 'user_email': str(user_email) if user_email else '', 'api_requests': str(row.get('api_requests', 0)), 'successful_requests': str(row.get('successful_requests', 0)), 'failed_requests': str(row.get('failed_requests', 0)), @@ -184,4 +186,3 @@ class CBFTransformer: return None - diff --git a/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py b/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py index 5ba6457f376..97daaa32557 100644 --- a/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py +++ b/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py @@ -32,6 +32,7 @@ class TestCloudZeroDryRunEndpoint: 'team_id': ['team1', 'team2'], 'team_alias': ['Team One', 'Team Two'], 'api_key_alias': ['key1', 'key2'], + 'user_email': ['one@example.com', None], 'prompt_tokens': [100, 200], 'completion_tokens': [50, 100], 'spend': [0.01, 0.02], @@ -51,7 +52,8 @@ class TestCloudZeroDryRunEndpoint: 'entity_id': ['team1', 'team2'], 'resource/tag:team_id': ['team1', 'team2'], 'resource/tag:team_alias': ['Team One', 'Team Two'], - 'resource/tag:api_key_alias': ['key1', 'key2'] + 'resource/tag:api_key_alias': ['key1', 'key2'], + 'resource/tag:user_email': ['one@example.com', 'N/A'] }) with patch('litellm.integrations.cloudzero.database.LiteLLMDatabase') as mock_db_class, \ @@ -86,6 +88,7 @@ class TestCloudZeroDryRunEndpoint: assert len(result['cbf_data']) == 2 assert result['cbf_data'][0]['cost/cost'] == 0.01 assert result['cbf_data'][1]['cost/cost'] == 0.02 + assert result['cbf_data'][0]['resource/tag:user_email'] == 'one@example.com' # Verify summary summary = result['summary'] @@ -122,4 +125,3 @@ class TestCloudZeroDryRunEndpoint: assert result['summary']['total_records'] == 0 assert result['summary']['total_cost'] == 0 assert result['summary']['total_tokens'] == 0 - diff --git a/tests/test_litellm/integrations/cloudzero/test_transform.py b/tests/test_litellm/integrations/cloudzero/test_transform.py index 1f4db10cab8..468f96ece1d 100644 --- a/tests/test_litellm/integrations/cloudzero/test_transform.py +++ b/tests/test_litellm/integrations/cloudzero/test_transform.py @@ -116,6 +116,56 @@ class TestCBFTransformer: assert result['usage/units'] == 'tokens' assert result['resource/id'] == 'test-czrn' + def test_create_cbf_record_adds_user_email_tag(self): + """Test that user_email field is emitted as a resource tag when present.""" + transformer = CBFTransformer() + with patch.object(transformer.czrn_generator, 'create_from_litellm_data') as mock_czrn, \ + patch.object(transformer.czrn_generator, 'extract_components') as mock_extract: + + mock_czrn.return_value = 'test-czrn' + mock_extract.return_value = ('service', 'provider', 'region', 'account', 'resource', 'local_id') + + row = { + 'date': '2025-01-19', + 'spend': 1.0, + 'prompt_tokens': 10, + 'completion_tokens': 5, + 'model': 'gpt-4', + 'api_key': 'sk-useremail', + 'team_id': 'team-123', + 'team_alias': 'Dev Team', + 'user_email': 'user@example.com' + } + + result = transformer._create_cbf_record(row) + + assert result['resource/tag:user_email'] == 'user@example.com' + + def test_create_cbf_record_omits_empty_user_email(self): + """Test that empty user_email values are not added as resource tags.""" + transformer = CBFTransformer() + with patch.object(transformer.czrn_generator, 'create_from_litellm_data') as mock_czrn, \ + patch.object(transformer.czrn_generator, 'extract_components') as mock_extract: + + mock_czrn.return_value = 'test-czrn' + mock_extract.return_value = ('service', 'provider', 'region', 'account', 'resource', 'local_id') + + row = { + 'date': '2025-01-19', + 'spend': 1.0, + 'prompt_tokens': 10, + 'completion_tokens': 5, + 'model': 'gpt-4', + 'api_key': 'sk-useremail', + 'team_id': 'team-123', + 'team_alias': 'Dev Team', + 'user_email': None + } + + result = transformer._create_cbf_record(row) + + assert 'resource/tag:user_email' not in result + def test_create_cbf_record_minimal_data(self): """Test _create_cbf_record method with minimal row data.""" transformer = CBFTransformer() @@ -180,4 +230,4 @@ class TestCBFTransformer: result = transformer._parse_date('2025-01-19T10:30:00Z') assert isinstance(result, datetime) - assert result.year == 2025 \ No newline at end of file + assert result.year == 2025 From 1112974112ecc4aacf7ca1ab811c7e3cbac48b8f Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 3 Jan 2026 19:22:38 -0800 Subject: [PATCH 015/304] Virtual Keys Table Loading State --- .../VirtualKeysPage/VirtualKeysTable.test.tsx | 132 +++++++++++++++++- .../VirtualKeysPage/VirtualKeysTable.tsx | 21 +-- 2 files changed, 142 insertions(+), 11 deletions(-) diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 3f55b11769c..cbd3d2c7320 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -1,4 +1,4 @@ -import { screen, waitFor } from "@testing-library/react"; +import { screen, waitFor, fireEvent } from "@testing-library/react"; import { vi, it, expect, beforeEach, MockedFunction } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; import { VirtualKeysTable } from "./VirtualKeysTable"; @@ -264,3 +264,133 @@ it("should show skeleton loaders when isLoading is true", () => { expect(screen.queryByText("Test Key Alias")).not.toBeInTheDocument(); expect(screen.queryByText("Test Team")).not.toBeInTheDocument(); }); + +it("should show 'No keys found' message when filteredKeys is empty", () => { + // Mock empty filteredKeys + mockUseFilterLogic.mockReturnValue({ + filters: { + "Team ID": "", + "Organization ID": "", + "Key Alias": "", + "User ID": "", + "Sort By": "created_at", + "Sort Order": "desc", + }, + filteredKeys: [], + allKeyAliases: [], + allTeams: [mockTeam], + allOrganizations: [mockOrganization], + handleFilterChange: vi.fn(), + handleFilterReset: vi.fn(), + }); + + const mockProps = { + teams: [mockTeam], + organizations: [mockOrganization], + onSortChange: vi.fn(), + currentSort: { + sortBy: "created_at", + sortOrder: "desc" as const, + }, + }; + + renderWithProviders(); + + expect(screen.getByText("No keys found")).toBeInTheDocument(); +}); + +it("should handle models with more than 3 entries to trigger expansion UI", () => { + const keyWithManyModels = { + ...mockKey, + models: ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "claude-3", "claude-3-5-sonnet"], + }; + + mockUseFilterLogic.mockReturnValue({ + filters: { + "Team ID": "", + "Organization ID": "", + "Key Alias": "", + "User ID": "", + "Sort By": "created_at", + "Sort Order": "desc", + }, + filteredKeys: [keyWithManyModels], + allKeyAliases: ["test-key-alias"], + allTeams: [mockTeam], + allOrganizations: [mockOrganization], + handleFilterChange: vi.fn(), + handleFilterReset: vi.fn(), + }); + + const mockProps = { + teams: [mockTeam], + organizations: [mockOrganization], + onSortChange: vi.fn(), + currentSort: { + sortBy: "created_at", + sortOrder: "desc" as const, + }, + }; + + renderWithProviders(); + + // This test ensures the ChevronDownIcon import (line 6) is used + // by having a key with > 3 models which triggers the expansion logic + // that uses ChevronDownIcon and ChevronRightIcon + expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); +}); + +it("should render table headers correctly", () => { + const mockProps = { + teams: [mockTeam], + organizations: [mockOrganization], + onSortChange: vi.fn(), + currentSort: { + sortBy: "created_at", + sortOrder: "desc" as const, + }, + }; + + renderWithProviders(); + + // Check that main headers are rendered (testing the header.isPlaceholder condition path) + expect(screen.getByText("Key ID")).toBeInTheDocument(); + expect(screen.getByText("Key Alias")).toBeInTheDocument(); + expect(screen.getByText("Team Alias")).toBeInTheDocument(); + expect(screen.getByText("Models")).toBeInTheDocument(); + expect(screen.getByText("Spend (USD)")).toBeInTheDocument(); +}); + +it("should handle column resizing hover events", () => { + const mockProps = { + teams: [mockTeam], + organizations: [mockOrganization], + onSortChange: vi.fn(), + currentSort: { + sortBy: "created_at", + sortOrder: "desc" as const, + }, + }; + + renderWithProviders(); + + // Find a header cell with data-header-id attribute + const headerCell = document.querySelector("[data-header-id]") as HTMLElement; + + expect(headerCell).toBeInTheDocument(); + + // Check that the resizer element exists within the header + const resizer = headerCell?.querySelector(".resizer") as HTMLElement; + expect(resizer).toBeInTheDocument(); + + // Initially, resizer should have opacity 0 + expect(resizer.style.opacity).toBe("0"); + + // Simulate mouse enter using fireEvent - should set opacity to 0.5 (lines 612-616) + fireEvent.mouseEnter(headerCell); + expect(resizer.style.opacity).toBe("0.5"); + + // Simulate mouse leave using fireEvent - should set opacity back to 0 (lines 618-622) + fireEvent.mouseLeave(headerCell); + expect(resizer.style.opacity).toBe("0"); +}); diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index b95d675979c..3bda8ee2f02 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -68,12 +68,13 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo }); const [tablePagination, setTablePagination] = React.useState({ pageIndex: 0, - pageSize: 100, + pageSize: 50, }); const { data: keys, isPending: isLoading, + isFetching, refetch, } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize); const totalCount = keys?.total_count || 0; @@ -545,8 +546,8 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
- {isLoading ? ( - + {isLoading || isFetching ? ( + ) : ( Showing {rangeLabel} of {totalCount} results @@ -554,32 +555,32 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo )}
- {isLoading ? ( - + {isLoading || isFetching ? ( + ) : ( Page {pageIndex + 1} of {table.getPageCount()} )} - {isLoading ? ( + {isLoading || isFetching ? ( ) : ( )} - {isLoading ? ( + {isLoading || isFetching ? ( ) : ( -
- {isFetching ? ( -
- Loading configuration... -
- ) : Object.keys(discountConfig).length > 0 ? ( - - ) : ( -
- + +
+ Provider Discounts + + Apply percentage-based discounts to reduce costs for specific providers + +
+ + + + + Discounts + Test It + + + +
+
+
- )} -
-
- -
- -
-
-
-
-
- + {isFetching ? ( +
+ Loading configuration... +
+ ) : Object.keys(discountConfig).length > 0 ? ( + + ) : ( +
+ + + + + No provider discounts configured + + + Click "Add Provider Discount" to get started + +
+ )} +
+ + +
+ +
+
+ + + + + )} - {/* Accordion 2: Fee/Price Margin */} - + {/* Accordion 2: Fee/Price Margin - Only for proxy admins */} + {isProxyAdmin && ( + + +
+ Fee/Price Margin + + Add fees or margins to LLM costs for internal billing and cost recovery + +
+
+ +
+
+ +
+ {isFetching ? ( +
+ Loading configuration... +
+ ) : Object.keys(marginConfig).length > 0 ? ( + + ) : ( +
+ + + + + No provider margins configured + + + Click "Add Provider Margin" to get started + +
+ )} +
+
+
+ )} + + {/* Accordion 3: Pricing Calculator - Available to all roles */} +
- Fee/Price Margin + Pricing Calculator - Add fees or margins to LLM costs for internal billing and cost recovery + Estimate LLM costs based on expected token usage and request volume
-
- -
- {isFetching ? ( -
- Loading configuration... -
- ) : Object.keys(marginConfig).length > 0 ? ( - - ) : ( -
- - - - - No provider margins configured - - - Click "Add Provider Margin" to get started - -
- )} +
diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/cost_results.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/cost_results.tsx new file mode 100644 index 00000000000..03d5ca0e518 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/cost_results.tsx @@ -0,0 +1,203 @@ +import React from "react"; +import { Text } from "@tremor/react"; +import { Card, Statistic, Row, Col, Divider, Spin } from "antd"; +import { DollarOutlined, LoadingOutlined } from "@ant-design/icons"; +import { CostEstimateResponse } from "../types"; +import { formatNumberWithCommas } from "@/utils/dataUtils"; +import ExportDropdown from "./export_dropdown"; + +interface CostResultsProps { + result: CostEstimateResponse | null; + loading: boolean; +} + +const formatCost = (value: number | null | undefined): string => { + if (value === null || value === undefined) return "-"; + if (value === 0) return "$0"; + if (value < 0.0001) return `$${value.toExponential(2)}`; + if (value < 1) return `$${value.toFixed(4)}`; + return `$${formatNumberWithCommas(value, 2, true)}`; +}; + +const formatRequests = (value: number | null | undefined): string => { + if (value === null || value === undefined) return "-"; + return formatNumberWithCommas(value, 0, true); +}; + +const CostResults: React.FC = ({ result, loading }) => { + if (!result && !loading) { + return ( +
+ + Select a model to see cost estimates + +
+ ); + } + + if (loading && !result) { + return ( +
+ } /> + Calculating costs... +
+ ); + } + + if (!result) return null; + + return ( +
+ + +
+
+ Cost Estimate + + Model: {result.model} {result.provider && `(${result.provider})`} + +
+
+ {loading && } size="small" />} + +
+
+ + + + + } + /> + + + + + + + + + 0 ? "#faad14" : undefined, + }} + /> + + + + + {result.daily_cost !== null && ( + + + + } + /> + + + + + + + + + 0 ? "#faad14" : undefined, + }} + /> + + + + )} + + {result.monthly_cost !== null && ( + + + + } + /> + + + + + + + + + 0 ? "#faad14" : undefined, + }} + /> + + + + )} + + {(result.input_cost_per_token || result.output_cost_per_token) && ( +
+ Token Pricing: + {result.input_cost_per_token && ( + Input: ${formatNumberWithCommas(result.input_cost_per_token * 1_000_000, 2)}/1M tokens + )} + {result.input_cost_per_token && result.output_cost_per_token && " | "} + {result.output_cost_per_token && ( + Output: ${formatNumberWithCommas(result.output_cost_per_token * 1_000_000, 2)}/1M tokens + )} +
+ )} +
+ ); +}; + +export default CostResults; diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_dropdown.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_dropdown.tsx new file mode 100644 index 00000000000..e8a681021d6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_dropdown.tsx @@ -0,0 +1,71 @@ +import React, { useState, useRef, useEffect } from "react"; +import { Button } from "@tremor/react"; +import { DownloadOutlined, FilePdfOutlined, FileExcelOutlined } from "@ant-design/icons"; +import { CostEstimateResponse } from "../types"; +import { exportToPDF, exportToCSV } from "./export_utils"; + +interface ExportDropdownProps { + result: CostEstimateResponse; +} + +const ExportDropdown: React.FC = ({ result }) => { + const [isOpen, setIsOpen] = useState(false); + const menuRef = useRef(null); + + useEffect(() => { + const handleClickOutside = (event: MouseEvent) => { + if (menuRef.current && !menuRef.current.contains(event.target as Node)) { + setIsOpen(false); + } + }; + + if (isOpen) { + document.addEventListener("mousedown", handleClickOutside); + } + + return () => { + document.removeEventListener("mousedown", handleClickOutside); + }; + }, [isOpen]); + + return ( +
+ + + {isOpen && ( +
+ + +
+ )} +
+ ); +}; + +export default ExportDropdown; + diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_utils.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_utils.ts new file mode 100644 index 00000000000..e02e8288456 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/export_utils.ts @@ -0,0 +1,276 @@ +import { CostEstimateResponse } from "../types"; +import { formatNumberWithCommas } from "@/utils/dataUtils"; + +const formatCostForExport = (value: number | null | undefined): string => { + if (value === null || value === undefined) return "-"; + if (value === 0) return "$0.00"; + if (value < 0.01) return `$${value.toFixed(6)}`; + if (value < 1) return `$${value.toFixed(4)}`; + return `$${formatNumberWithCommas(value, 2)}`; +}; + +const formatRequestsForExport = (value: number | null | undefined): string => { + if (value === null || value === undefined) return "-"; + return formatNumberWithCommas(value, 0); +}; + +export const exportToPDF = (result: CostEstimateResponse): void => { + const printWindow = window.open("", "_blank"); + if (!printWindow) { + alert("Please allow popups to export PDF"); + return; + } + + const html = ` + + + + Cost Estimate Report - ${result.model} + + + +

LLM Cost Estimate Report

+ +
+

Model: ${result.model}

+ ${result.provider ? `

Provider: ${result.provider}

` : ""} +

Input Tokens per Request: ${formatRequestsForExport(result.input_tokens)}

+

Output Tokens per Request: ${formatRequestsForExport(result.output_tokens)}

+ ${result.num_requests_per_day ? `

Requests per Day: ${formatRequestsForExport(result.num_requests_per_day)}

` : ""} + ${result.num_requests_per_month ? `

Requests per Month: ${formatRequestsForExport(result.num_requests_per_month)}

` : ""} +
+ +

Per-Request Cost Breakdown

+ + + + + + + + + + + + + + + + + + + + + +
Cost TypeAmount
Input Cost${formatCostForExport(result.input_cost_per_request)}
Output Cost${formatCostForExport(result.output_cost_per_request)}
Margin/Fee${formatCostForExport(result.margin_cost_per_request)}
Total per Request${formatCostForExport(result.cost_per_request)}
+ + ${result.daily_cost !== null ? ` +

Daily Cost Estimate (${formatRequestsForExport(result.num_requests_per_day)} requests/day)

+ + + + + + + + + + + + + + + + + + + + + +
Cost TypeAmount
Input Cost${formatCostForExport(result.daily_input_cost)}
Output Cost${formatCostForExport(result.daily_output_cost)}
Margin/Fee${formatCostForExport(result.daily_margin_cost)}
Total Daily${formatCostForExport(result.daily_cost)}
+ ` : ""} + + ${result.monthly_cost !== null ? ` +

Monthly Cost Estimate (${formatRequestsForExport(result.num_requests_per_month)} requests/month)

+ + + + + + + + + + + + + + + + + + + + + +
Cost TypeAmount
Input Cost${formatCostForExport(result.monthly_input_cost)}
Output Cost${formatCostForExport(result.monthly_output_cost)}
Margin/Fee${formatCostForExport(result.monthly_margin_cost)}
Total Monthly${formatCostForExport(result.monthly_cost)}
+ ` : ""} + + ${result.input_cost_per_token || result.output_cost_per_token ? ` +

Token Pricing

+ + + + + + ${result.input_cost_per_token ? ` + + + + + ` : ""} + ${result.output_cost_per_token ? ` + + + + + ` : ""} +
Token TypePrice per 1M Tokens
Input Tokens$${(result.input_cost_per_token * 1000000).toFixed(2)}
Output Tokens$${(result.output_cost_per_token * 1000000).toFixed(2)}
+ ` : ""} + + + + + `; + + printWindow.document.write(html); + printWindow.document.close(); + printWindow.onload = () => { + printWindow.print(); + }; +}; + +export const exportToCSV = (result: CostEstimateResponse): void => { + const rows = [ + ["LLM Cost Estimate Report"], + [""], + ["Configuration"], + ["Model", result.model], + ["Provider", result.provider || "-"], + ["Input Tokens per Request", result.input_tokens.toString()], + ["Output Tokens per Request", result.output_tokens.toString()], + ["Requests per Day", result.num_requests_per_day?.toString() || "-"], + ["Requests per Month", result.num_requests_per_month?.toString() || "-"], + [""], + ["Per-Request Costs"], + ["Input Cost", result.input_cost_per_request.toString()], + ["Output Cost", result.output_cost_per_request.toString()], + ["Margin/Fee", result.margin_cost_per_request.toString()], + ["Total per Request", result.cost_per_request.toString()], + ]; + + if (result.daily_cost !== null) { + rows.push( + [""], + ["Daily Costs"], + ["Daily Input Cost", result.daily_input_cost?.toString() || "-"], + ["Daily Output Cost", result.daily_output_cost?.toString() || "-"], + ["Daily Margin/Fee", result.daily_margin_cost?.toString() || "-"], + ["Total Daily", result.daily_cost.toString()] + ); + } + + if (result.monthly_cost !== null) { + rows.push( + [""], + ["Monthly Costs"], + ["Monthly Input Cost", result.monthly_input_cost?.toString() || "-"], + ["Monthly Output Cost", result.monthly_output_cost?.toString() || "-"], + ["Monthly Margin/Fee", result.monthly_margin_cost?.toString() || "-"], + ["Total Monthly", result.monthly_cost.toString()] + ); + } + + if (result.input_cost_per_token || result.output_cost_per_token) { + rows.push( + [""], + ["Token Pricing (per 1M tokens)"], + ["Input Token Price", result.input_cost_per_token ? `$${(result.input_cost_per_token * 1000000).toFixed(2)}` : "-"], + ["Output Token Price", result.output_cost_per_token ? `$${(result.output_cost_per_token * 1000000).toFixed(2)}` : "-"] + ); + } + + const csv = rows.map(row => row.join(",")).join("\n"); + const blob = new Blob([csv], { type: "text/csv;charset=utf-8;" }); + const url = window.URL.createObjectURL(blob); + const a = document.createElement("a"); + a.href = url; + a.download = `cost_estimate_${result.model.replace(/\//g, "_")}_${new Date().toISOString().split("T")[0]}.csv`; + document.body.appendChild(a); + a.click(); + document.body.removeChild(a); + window.URL.revokeObjectURL(url); +}; + diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/index.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/index.tsx new file mode 100644 index 00000000000..fff6475e8a2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/index.tsx @@ -0,0 +1,31 @@ +import React, { useCallback } from "react"; +import PricingForm from "./pricing_form"; +import CostResults from "./cost_results"; +import { useCostEstimate } from "./use_cost_estimate"; +import { PricingCalculatorProps, PricingFormValues } from "./types"; + +const PricingCalculator: React.FC = ({ + accessToken, + models, +}) => { + const { loading, result, debouncedFetch } = useCostEstimate(accessToken); + + const handleValuesChange = useCallback( + (_changedValues: Partial, allValues: PricingFormValues) => { + if (allValues.model) { + debouncedFetch(allValues); + } + }, + [debouncedFetch] + ); + + return ( +
+ + +
+ ); +}; + +export default PricingCalculator; + diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/pricing_form.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/pricing_form.tsx new file mode 100644 index 00000000000..c91a19e5516 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/pricing_form.tsx @@ -0,0 +1,104 @@ +import React from "react"; +import { Form, InputNumber, Select, Row, Col } from "antd"; +import { PricingFormValues } from "./types"; + +interface PricingFormProps { + models: string[]; + onValuesChange: (changedValues: Partial, allValues: PricingFormValues) => void; +} + +const PricingForm: React.FC = ({ models, onValuesChange }) => { + return ( +
+ + + + handleEntryChange(record.id, "model", value)} + optionFilterProp="label" + filterOption={(input, option) => + String(option?.label ?? "").toLowerCase().includes(input.toLowerCase()) + } + options={models.map((model) => ({ + value: model, + label: model, + }))} + style={{ width: "100%" }} + size="small" + /> + ), + }, + { + title: "Input Tokens", + dataIndex: "input_tokens", + key: "input_tokens", + width: "18%", + render: (_: number, record: ModelEntry) => ( + handleEntryChange(record.id, "input_tokens", value ?? 0)} + style={{ width: "100%" }} + size="small" + formatter={(value) => `${value}`.replace(/\B(?=(\d{3})+(?!\d))/g, ",")} + /> + ), + }, + { + title: "Output Tokens", + dataIndex: "output_tokens", + key: "output_tokens", + width: "18%", + render: (_: number, record: ModelEntry) => ( + handleEntryChange(record.id, "output_tokens", value ?? 0)} + style={{ width: "100%" }} + size="small" + formatter={(value) => `${value}`.replace(/\B(?=(\d{3})+(?!\d))/g, ",")} + /> + ), + }, + { + title: `Requests/${timePeriod === "day" ? "Day" : "Month"}`, + dataIndex: timePeriod === "day" ? "num_requests_per_day" : "num_requests_per_month", + key: "num_requests", + width: "20%", + render: (_: number | undefined, record: ModelEntry) => ( + + handleEntryChange( + record.id, + timePeriod === "day" ? "num_requests_per_day" : "num_requests_per_month", + value ?? undefined + ) + } + style={{ width: "100%" }} + size="small" + placeholder="-" + formatter={(value) => (value ? `${value}`.replace(/\B(?=(\d{3})+(?!\d))/g, ",") : "")} + /> + ), + }, + { + title: "", + key: "actions", + width: 50, + render: (_: unknown, record: ModelEntry) => ( +