From ae867e0eac92e72b9a997f44eb8b710be6b4fb1f Mon Sep 17 00:00:00 2001 From: mutnale_sushant Date: Thu, 9 Jul 2026 07:02:14 +0530 Subject: [PATCH] feat(param_utils): prevent internal parameter leakage and fix logger registry typing I implemented a new strip_litellm_internal_params utility to recursively filter out internal LiteLLM parameters (such as litellm_params, litellm_metadata, proxy_server_request, and any custom parameter starting with _litellm_) from request payloads and nested extra_body structures. I applied this sanitization logic to OpenAI, Azure, and OpenAI-like chat completions and embeddings before request execution. I fixed a type checking mismatch in the custom logger registry by type annotating CALLBACK_CLASS_STR_TO_CLASS_TYPE as dict[str, type[object]] to support dynamic registration of enterprise loggers. I added formal pytest unit tests in tests/test_litellm/test_openai_params_strip.py covering sync/async completions and embeddings. Fixes #14901 --- .../custom_logger_registry.py | 2 +- litellm/litellm_core_utils/param_utils.py | 45 +++++ litellm/llms/azure/azure.py | 14 +- litellm/llms/openai/openai.py | 16 +- litellm/llms/openai_like/chat/handler.py | 13 +- .../test_litellm/test_openai_params_strip.py | 158 ++++++++++++++++++ 6 files changed, 234 insertions(+), 14 deletions(-) create mode 100644 litellm/litellm_core_utils/param_utils.py create mode 100644 tests/test_litellm/test_openai_params_strip.py diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index a7fae104c92..eb7bbc79971 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -60,7 +60,7 @@ class CustomLoggerRegistry: Registry mapping the callback class string to the class type. """ - CALLBACK_CLASS_STR_TO_CLASS_TYPE = { + CALLBACK_CLASS_STR_TO_CLASS_TYPE: dict[str, type[object]] = { "lago": LagoLogger, "openmeter": OpenMeterLogger, "braintrust": BraintrustLogger, diff --git a/litellm/litellm_core_utils/param_utils.py b/litellm/litellm_core_utils/param_utils.py new file mode 100644 index 00000000000..d3eeda00c60 --- /dev/null +++ b/litellm/litellm_core_utils/param_utils.py @@ -0,0 +1,45 @@ +import logging + +logger = logging.getLogger(__name__) + +LITELLM_INTERNAL_PARAM_NAMES = { + "litellm_params", + "proxy_server_request", + "model_info", + "metadata", + "preset_cache_key", + "litellm_metadata", + "acompletion", +} + + +def strip_litellm_internal_params(data: dict[str, object]) -> dict[str, object]: + """ + Remove LiteLLM internal params (e.g. litellm_params, proxy_server_request, _litellm_ prefixed keys) + from request data before passing to client libraries (e.g. OpenAI). + + This avoids throwing API validation/schema errors (e.g. 400 Bad Request) due to unknown parameters. + """ + if not isinstance(data, dict): # pyright: ignore[reportUnnecessaryIsInstance] # runtime guard for unsanitized input + return data # pyright: ignore[reportUnreachable] # runtime guard + + try: + # Create a shallow copy so we don't modify the input dictionary in-place + cleaned_data: dict[str, object] = {} + for key, value in data.items(): + if key in LITELLM_INTERNAL_PARAM_NAMES or key.startswith("_litellm_"): + continue + if key == "extra_body" and isinstance(value, dict): + cleaned_extra_body: dict[str, object] = {} + extra_body_dict: dict[object, object] = value # pyright: ignore[reportUnknownVariableType] # cast from dynamic dict + for k, v in extra_body_dict.items(): + if isinstance(k, str) and (k in LITELLM_INTERNAL_PARAM_NAMES or k.startswith("_litellm_")): + continue + cleaned_extra_body[str(k)] = v + cleaned_data["extra_body"] = cleaned_extra_body + else: + cleaned_data[key] = value + return cleaned_data + except Exception as e: + logger.warning(f"Error in strip_litellm_internal_params: {str(e)}") + return data diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index ccb9eb8f5c8..67efac19499 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -33,6 +33,8 @@ from litellm.utils import ( convert_to_model_response_object, modify_url, ) +from litellm.litellm_core_utils.param_utils import strip_litellm_internal_params + from ...types.llms.openai import HttpxBinaryResponseContent from ..base import BaseLLM @@ -145,7 +147,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): - call chat.completions.create by default """ try: - raw_response = azure_client.chat.completions.with_raw_response.create(**data, timeout=timeout) + cleaned_data = strip_litellm_internal_params(data) + raw_response = azure_client.chat.completions.with_raw_response.create(**cleaned_data, timeout=timeout) headers = dict(raw_response.headers) response = raw_response.parse() @@ -168,7 +171,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): """ start_time = time.time() try: - raw_response = await azure_client.chat.completions.with_raw_response.create(**data, timeout=timeout) + cleaned_data = strip_litellm_internal_params(data) + raw_response = await azure_client.chat.completions.with_raw_response.create(**cleaned_data, timeout=timeout) headers = dict(raw_response.headers) response = raw_response.parse() @@ -667,7 +671,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): if not isinstance(openai_aclient, (AsyncAzureOpenAI, AsyncOpenAI)): raise ValueError("Azure client is not an instance of AsyncAzureOpenAI or AsyncOpenAI") - raw_response = await openai_aclient.embeddings.with_raw_response.create(**data, timeout=timeout) + cleaned_data = strip_litellm_internal_params(data) + raw_response = await openai_aclient.embeddings.with_raw_response.create(**cleaned_data, timeout=timeout) headers = dict(raw_response.headers) # Convert json.JSONDecodeError to AzureOpenAIError for two critical reasons: @@ -793,7 +798,8 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ) ## COMPLETION CALL - raw_response = azure_client.embeddings.with_raw_response.create(**data, timeout=timeout) # type: ignore + cleaned_data = strip_litellm_internal_params(data) + raw_response = azure_client.embeddings.with_raw_response.create(**cleaned_data, timeout=timeout) # type: ignore headers = dict(raw_response.headers) response = raw_response.parse() if isinstance(response, str): diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 6b191144a11..8048e13068b 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -50,6 +50,8 @@ from litellm.utils import ( ProviderConfigManager, convert_to_model_response_object, ) +from litellm.litellm_core_utils.param_utils import strip_litellm_internal_params + from ...types.llms.openai import * from ..base import BaseLLM @@ -424,7 +426,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): """ start_time = time.time() try: - raw_response = await openai_aclient.chat.completions.with_raw_response.create(**data, timeout=timeout) + cleaned_data = strip_litellm_internal_params(data) + raw_response = await openai_aclient.chat.completions.with_raw_response.create( + **cleaned_data, timeout=timeout + ) end_time = time.time() if hasattr(raw_response, "headers"): @@ -461,7 +466,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): """ raw_response = None try: - raw_response = openai_client.chat.completions.with_raw_response.create(**data, timeout=timeout) + cleaned_data = strip_litellm_internal_params(data) + raw_response = openai_client.chat.completions.with_raw_response.create(**cleaned_data, timeout=timeout) if hasattr(raw_response, "headers"): headers = dict(raw_response.headers) @@ -1153,7 +1159,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): - call embeddings.create by default """ try: - raw_response = await openai_aclient.embeddings.with_raw_response.create(**data, timeout=timeout) # type: ignore + cleaned_data = strip_litellm_internal_params(data) + raw_response = await openai_aclient.embeddings.with_raw_response.create(**cleaned_data, timeout=timeout) # type: ignore headers = dict(raw_response.headers) response = raw_response.parse() return headers, response @@ -1174,7 +1181,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): - call embeddings.create by default """ try: - raw_response = openai_client.embeddings.with_raw_response.create(**data, timeout=timeout) # type: ignore + cleaned_data = strip_litellm_internal_params(data) + raw_response = openai_client.embeddings.with_raw_response.create(**cleaned_data, timeout=timeout) # type: ignore headers = dict(raw_response.headers) response = raw_response.parse() diff --git a/litellm/llms/openai_like/chat/handler.py b/litellm/llms/openai_like/chat/handler.py index 0da0f3d90f0..11d798a9aa2 100644 --- a/litellm/llms/openai_like/chat/handler.py +++ b/litellm/llms/openai_like/chat/handler.py @@ -18,6 +18,8 @@ from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.openai.openai import OpenAIConfig from litellm.types.utils import CustomStreamingDecoder, ModelResponse from litellm.utils import CustomStreamWrapper, ProviderConfigManager +from litellm.litellm_core_utils.param_utils import strip_litellm_internal_params + from ..common_utils import OpenAILikeBase, OpenAILikeError from .transformation import OpenAILikeChatConfig @@ -268,17 +270,18 @@ class OpenAILikeChatHandler(OpenAILikeBase): "headers": headers, }, ) + cleaned_data = strip_litellm_internal_params(data) if acompletion is True: if client is None or not isinstance(client, AsyncHTTPHandler): client = None if ( stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) - data["stream"] = stream + cleaned_data["stream"] = stream return self.acompletion_stream_function( model=model, messages=messages, - data=data, + data=cleaned_data, api_base=api_base, custom_prompt_dict=custom_prompt_dict, model_response=model_response, @@ -300,7 +303,7 @@ class OpenAILikeChatHandler(OpenAILikeBase): return self.acompletion_function( model=model, messages=messages, - data=data, + data=cleaned_data, api_base=api_base, custom_prompt_dict=custom_prompt_dict, custom_llm_provider=custom_llm_provider, @@ -326,7 +329,7 @@ class OpenAILikeChatHandler(OpenAILikeBase): client=(client if client is not None and isinstance(client, HTTPHandler) else None), api_base=api_base, headers=headers, - data=json.dumps(data), + data=json.dumps(cleaned_data), model=model, messages=messages, logging_obj=logging_obj, @@ -345,7 +348,7 @@ class OpenAILikeChatHandler(OpenAILikeBase): if client is None or not isinstance(client, HTTPHandler): client = HTTPHandler(timeout=timeout) # type: ignore try: - response = client.post(url=api_base, headers=headers, data=json.dumps(data)) + response = client.post(url=api_base, headers=headers, data=json.dumps(cleaned_data)) response.raise_for_status() except httpx.HTTPStatusError as e: diff --git a/tests/test_litellm/test_openai_params_strip.py b/tests/test_litellm/test_openai_params_strip.py new file mode 100644 index 00000000000..f145b022f0d --- /dev/null +++ b/tests/test_litellm/test_openai_params_strip.py @@ -0,0 +1,158 @@ +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +# Add project root to sys.path +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) + +import pytest + +import litellm +from litellm import acompletion, completion, embedding + +litellm.return_response_headers = False + + +@pytest.mark.asyncio +async def test_openai_chat_completion_params_strip(): + """ + Test that litellm_params and _litellm_* prefixed params are stripped + from OpenAI completion calls. + """ + # Mock return value of parse() which is what is called on raw_response + mock_choice = MagicMock() + mock_choice.finish_reason = "stop" + mock_choice.index = 0 + mock_choice.message = MagicMock(content="Mock response", role="assistant") + mock_choice.message.tool_calls = None + mock_choice.message.function_call = None + mock_choice.message.provider_specific_fields = {} + + mock_response_data = MagicMock() + mock_response_data.choices = [mock_choice] + mock_response_data.id = "chatcmpl-123" + mock_response_data.created = 1677858242 + mock_response_data.model = "gpt-4o" + mock_response_data.object = "chat.completion" + mock_response_data.usage = MagicMock(completion_tokens=10, prompt_tokens=5, total_tokens=15) + + # We mock the underlying client create call + mock_create = MagicMock() + mock_raw_resp = MagicMock() + mock_raw_resp.headers = {"x-test-header": "test"} + mock_raw_resp.parse.return_value = mock_response_data + mock_create.return_value = mock_raw_resp + + with patch("openai.resources.chat.completions.Completions.create", mock_create): + completion( + model="gpt-4o", + messages=[{"role": "user", "content": "hi"}], + api_key="mock-key", + # internal params that should be stripped + litellm_params={"metadata": {"some_internal_key": "some_value"}}, + _litellm_test_param="test_value", + ) + + # Verify call arguments + mock_create.assert_called_once() + call_kwargs = mock_create.call_args[1] + + # Verify that internal params are not in the top-level keys or extra_body + assert "litellm_params" not in call_kwargs + assert "_litellm_test_param" not in call_kwargs + + extra_body = call_kwargs.get("extra_body", {}) + if extra_body: + assert "litellm_params" not in extra_body + assert "_litellm_test_param" not in extra_body + + +@pytest.mark.asyncio +async def test_openai_chat_acompletion_params_strip(): + """ + Test that litellm_params and _litellm_* prefixed params are stripped + from OpenAI async completion calls. + """ + mock_choice = MagicMock() + mock_choice.finish_reason = "stop" + mock_choice.index = 0 + mock_choice.message = MagicMock(content="Mock response", role="assistant") + mock_choice.message.tool_calls = None + mock_choice.message.function_call = None + mock_choice.message.provider_specific_fields = {} + + mock_response_data = MagicMock() + mock_response_data.choices = [mock_choice] + mock_response_data.id = "chatcmpl-123" + mock_response_data.created = 1677858242 + mock_response_data.model = "gpt-4o" + mock_response_data.object = "chat.completion" + mock_response_data.usage = MagicMock(completion_tokens=10, prompt_tokens=5, total_tokens=15) + + mock_raw_resp = MagicMock() + mock_raw_resp.headers = {"x-test-header": "test"} + mock_raw_resp.parse.return_value = mock_response_data + + mock_acreate = AsyncMock(return_value=mock_raw_resp) + + with patch("openai.resources.chat.completions.AsyncCompletions.create", mock_acreate): + try: + await acompletion( + model="gpt-4o", + messages=[{"role": "user", "content": "hi"}], + api_key="mock-key", + litellm_params={"metadata": {"some_internal_key": "some_value"}}, + _litellm_test_param="test_value", + ) + except Exception: + pass + + mock_acreate.assert_called_once() + call_kwargs = mock_acreate.call_args[1] + + assert "litellm_params" not in call_kwargs + assert "_litellm_test_param" not in call_kwargs + + extra_body = call_kwargs.get("extra_body", {}) + if extra_body: + assert "litellm_params" not in extra_body + assert "_litellm_test_param" not in extra_body + + +@pytest.mark.asyncio +async def test_openai_embedding_params_strip(): + """ + Test that litellm_params and _litellm_* prefixed params are stripped + from OpenAI embedding calls. + """ + mock_response_data = MagicMock() + mock_response_data.model = "text-embedding-3-small" + mock_response_data.object = "list" + mock_response_data.data = [MagicMock(embedding=[0.1, 0.2])] + mock_response_data.usage = MagicMock(prompt_tokens=5, total_tokens=5) + + mock_create = MagicMock() + mock_raw_resp = MagicMock() + mock_raw_resp.headers = {"x-test-header": "test"} + mock_raw_resp.parse.return_value = mock_response_data + mock_create.return_value = mock_raw_resp + + with patch("openai.resources.embeddings.Embeddings.create", mock_create): + embedding( + model="text-embedding-3-small", + input=["hello"], + api_key="mock-key", + litellm_params={"metadata": {"some_internal_key": "some_value"}}, + _litellm_test_param="test_value", + ) + + mock_create.assert_called_once() + call_kwargs = mock_create.call_args[1] + + assert "litellm_params" not in call_kwargs + assert "_litellm_test_param" not in call_kwargs + + extra_body = call_kwargs.get("extra_body", {}) + if extra_body: + assert "litellm_params" not in extra_body + assert "_litellm_test_param" not in extra_body