diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index e03c4c93bd7..551f870aca9 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -1,5 +1,5 @@ """ -Support for gpt model family +Support for gpt model family """ from typing import ( @@ -11,6 +11,7 @@ from typing import ( List, Literal, Optional, + Tuple, Union, cast, overload, @@ -56,6 +57,7 @@ from ..common_utils import OpenAIError if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.types.llms.openai import ChatCompletionToolParam LiteLLMLoggingObj = _LiteLLMLoggingObj else: @@ -318,10 +320,12 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): content_item = content_item_typed return content_item + # fmt: off + @overload def _transform_messages( self, messages: List[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ) -> Coroutine[Any, Any, List[AllMessageValues]]: ... @overload @@ -333,6 +337,8 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ) -> List[AllMessageValues]: ... + # fmt: on + def _transform_messages( self, messages: List[AllMessageValues], model: str, is_async: bool = False ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: @@ -351,10 +357,10 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): List[OpenAIMessageContentListBlock], message_content ) for i, content_item in enumerate(message_content_types): - message_content_types[ - i - ] = await self._async_transform_content_item( - cast(OpenAIMessageContentListBlock, content_item), + message_content_types[i] = ( + await self._async_transform_content_item( + cast(OpenAIMessageContentListBlock, content_item), + ) ) return messages @@ -378,6 +384,29 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ) return messages + def remove_cache_control_flag_from_messages_and_tools( + self, + model: str, # allows overrides to selectively run this + messages: List[AllMessageValues], + tools: Optional[List["ChatCompletionToolParam"]] = None, + ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + filter_value_from_dict, + ) + from litellm.types.llms.openai import ChatCompletionToolParam + + for message in messages: + message = cast( + AllMessageValues, filter_value_from_dict(message, "cache_control") # type: ignore + ) + if tools is not None: + for tool in tools: + tool = cast( + ChatCompletionToolParam, + filter_value_from_dict(tool, "cache_control"), # type: ignore + ) + return messages, tools + def transform_request( self, model: str, @@ -393,6 +422,12 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): dict: The transformed request. Sent as the body of the API call. """ messages = self._transform_messages(messages=messages, model=model) + messages, tools = self.remove_cache_control_flag_from_messages_and_tools( + model=model, messages=messages, tools=optional_params.get("tools", []) + ) + if tools is not None and len(tools) > 0: + optional_params["tools"] = tools + return { "model": model, "messages": messages, @@ -410,7 +445,15 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): transformed_messages = await self._transform_messages( messages=messages, model=model, is_async=True ) - + transformed_messages, tools = ( + self.remove_cache_control_flag_from_messages_and_tools( + model=model, + messages=transformed_messages, + tools=optional_params.get("tools", []), + ) + ) + if tools is not None and len(tools) > 0: + optional_params["tools"] = tools if self.__class__._is_base_class: return { "model": model, diff --git a/litellm/llms/openrouter/chat/transformation.py b/litellm/llms/openrouter/chat/transformation.py index e3f9d5c3dd0..bf57218c91d 100644 --- a/litellm/llms/openrouter/chat/transformation.py +++ b/litellm/llms/openrouter/chat/transformation.py @@ -6,13 +6,13 @@ Calls done in OpenAI/openai.py as OpenRouter is openai-compatible. Docs: https://openrouter.ai/docs/parameters """ -from typing import Any, AsyncIterator, Iterator, List, Optional, Union +from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union import httpx from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.types.llms.openrouter import OpenRouterErrorMessage from litellm.types.utils import ModelResponse, ModelResponseStream @@ -43,11 +43,24 @@ class OpenrouterConfig(OpenAIGPTConfig): extra_body["models"] = models if route is not None: extra_body["route"] = route - mapped_openai_params[ - "extra_body" - ] = extra_body # openai client supports `extra_body` param + mapped_openai_params["extra_body"] = ( + extra_body # openai client supports `extra_body` param + ) return mapped_openai_params + def remove_cache_control_flag_from_messages_and_tools( + self, + model: str, + messages: List[AllMessageValues], + tools: Optional[List["ChatCompletionToolParam"]] = None, + ) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]: + if "claude" in model.lower(): # don't remove 'cache_control' flag + return messages, tools + else: + return super().remove_cache_control_flag_from_messages_and_tools( + model, messages, tools + ) + def transform_request( self, model: str, diff --git a/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py index 7272f8c8182..02eb872abab 100644 --- a/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py +++ b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py @@ -97,3 +97,20 @@ def test_openrouter_extra_body_transformation(): assert transformed_request["messages"] == [ {"role": "user", "content": "Hello, world!"} ] + + +def test_openrouter_cache_control_flag_removal(): + transformed_request = OpenrouterConfig().transform_request( + model="openrouter/deepseek/deepseek-chat", + messages=[ + { + "role": "user", + "content": "Hello, world!", + "cache_control": {"type": "ephemeral"}, + } + ], + optional_params={}, + litellm_params={}, + headers={}, + ) + assert transformed_request["messages"][0].get("cache_control") is None