mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Openrouter - filter out cache_control flag for non-anthropic models (allows usage with claude code) (#12850)
* fix(gpt_transformation.py): remove 'cache_control' flag for openai/openai-compatible calls Fixes https://github.com/BerriAI/litellm/issues/12787 * fix(openrouter/chat/transformation.py): allow passing openrouter cache control flag for claude models * fix(gpt_transformation.py): fix import * fix: fix adding tools
This commit is contained in:
parent
e4e10aa4ed
commit
e5251e7188
3 changed files with 85 additions and 12 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue