mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(openai): preserve cache_control for openai-compatible custom endpoints (#30387)
* fix(openai): preserve cache_control for openai-compatible custom endpoints * fix(openai): use parsed hostname to detect real OpenAI for cache_control preservation
This commit is contained in:
parent
8985913f87
commit
2a6b95a61f
2 changed files with 210 additions and 15 deletions
|
|
@ -18,6 +18,9 @@ from typing import (
|
|||
overload,
|
||||
)
|
||||
|
||||
import os
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
|
|
@ -426,6 +429,32 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
)
|
||||
return messages, tools
|
||||
|
||||
def _should_preserve_cache_control_for_endpoint(
|
||||
self,
|
||||
custom_llm_provider: Optional[str],
|
||||
api_base: Optional[str],
|
||||
) -> bool:
|
||||
"""
|
||||
The generic `openai` provider also reaches OpenAI-compatible endpoints
|
||||
(a LiteLLM proxy, vLLM, an Anthropic-compatible gateway) via a custom
|
||||
api_base. Those can understand cache_control, so it must survive there.
|
||||
Real OpenAI cannot, so it is still stripped for an openai.com host.
|
||||
"""
|
||||
if custom_llm_provider != "openai":
|
||||
return False
|
||||
resolved_api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or os.getenv("OPENAI_BASE_URL")
|
||||
or os.getenv("OPENAI_API_BASE")
|
||||
)
|
||||
if not resolved_api_base:
|
||||
return False
|
||||
hostname = urlparse(resolved_api_base).hostname
|
||||
if hostname is None:
|
||||
return False
|
||||
return hostname != "openai.com" and not hostname.endswith(".openai.com")
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -441,11 +470,14 @@ 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
|
||||
if not self._should_preserve_cache_control_for_endpoint(
|
||||
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
|
||||
):
|
||||
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
|
||||
|
||||
optional_params.pop("max_retries", None)
|
||||
|
||||
|
|
@ -466,16 +498,19 @@ 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 not self._should_preserve_cache_control_for_endpoint(
|
||||
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
|
||||
):
|
||||
(
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import pytest
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
|
||||
from litellm.llms.openai.chat.gpt_transformation import (
|
||||
OpenAIChatCompletionStreamingHandler,
|
||||
|
|
@ -571,3 +572,162 @@ class TestGPT5ReasoningEffortPreservation:
|
|||
|
||||
assert optional_params.get("temperature") == 0.5
|
||||
assert non_default_params.get("reasoning_effort") == "none"
|
||||
|
||||
|
||||
class TestCacheControlPreservationForCustomEndpoint:
|
||||
"""
|
||||
Regression tests for https://github.com/BerriAI/litellm/issues/30319
|
||||
|
||||
The AnthropicCacheControlHook injects cache_control when a user passes
|
||||
cache_control_injection_points, but the base OpenAIGPTConfig used to strip
|
||||
it unconditionally, making the feature a guaranteed no-op for the generic
|
||||
openai provider pointed at a cache_control-aware endpoint (a LiteLLM proxy,
|
||||
vLLM, an Anthropic-compatible gateway). cache_control must survive there
|
||||
while still being stripped for real api.openai.com.
|
||||
"""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = OpenAIGPTConfig()
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_openai_base_env(self, monkeypatch):
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
|
||||
monkeypatch.setattr(litellm, "api_base", None, raising=False)
|
||||
|
||||
@staticmethod
|
||||
def _cache_controlled_messages():
|
||||
return [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are helpful.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
]
|
||||
|
||||
def _transform(self, custom_llm_provider, api_base, optional_params=None):
|
||||
return self.config.transform_request(
|
||||
model="claude-sonnet-4",
|
||||
messages=self._cache_controlled_messages(),
|
||||
optional_params=optional_params or {},
|
||||
litellm_params={
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"api_base": api_base,
|
||||
},
|
||||
headers={},
|
||||
)
|
||||
|
||||
def test_predicate_openai_provider_custom_api_base_preserves(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint(
|
||||
"openai", "http://localhost:4000/v1"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_predicate_real_openai_no_api_base_strips(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint("openai", None)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_predicate_explicit_openai_host_strips(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint(
|
||||
"openai", "https://api.openai.com/v1"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_predicate_non_openai_provider_strips(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint(
|
||||
"deepseek", "https://api.deepseek.com"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_predicate_resolves_openai_base_url_env(self, monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_BASE_URL", "http://localhost:4000/v1")
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint("openai", None)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_predicate_resolves_openai_api_base_env(self, monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_API_BASE", "http://localhost:4000/v1")
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint("openai", None)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_predicate_lookalike_host_is_not_treated_as_openai(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint(
|
||||
"openai", "https://api.openai.com.evil.example/v1"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_predicate_openai_subdomain_strips(self):
|
||||
assert (
|
||||
self.config._should_preserve_cache_control_for_endpoint(
|
||||
"openai", "https://eu.api.openai.com/v1"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_transform_request_preserves_for_custom_api_base(self):
|
||||
body = self._transform("openai", "http://localhost:4000/v1")
|
||||
assert all("cache_control" in m for m in body["messages"])
|
||||
|
||||
def test_transform_request_strips_for_real_openai(self):
|
||||
body = self._transform("openai", None)
|
||||
assert all("cache_control" not in m for m in body["messages"])
|
||||
|
||||
def test_transform_request_strips_for_non_openai_provider(self):
|
||||
body = self._transform("fireworks_ai", "https://api.fireworks.ai/inference/v1")
|
||||
assert all("cache_control" not in m for m in body["messages"])
|
||||
|
||||
def test_transform_request_preserves_tool_cache_control(self):
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "f", "parameters": {}},
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
body = self._transform(
|
||||
"openai", "http://localhost:4000/v1", optional_params={"tools": tools}
|
||||
)
|
||||
assert "cache_control" in body["tools"][0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_transform_request_preserves_for_custom_api_base(self):
|
||||
body = await self.config.async_transform_request(
|
||||
model="claude-sonnet-4",
|
||||
messages=self._cache_controlled_messages(),
|
||||
optional_params={},
|
||||
litellm_params={
|
||||
"custom_llm_provider": "openai",
|
||||
"api_base": "http://localhost:4000/v1",
|
||||
},
|
||||
headers={},
|
||||
)
|
||||
assert all("cache_control" in m for m in body["messages"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_transform_request_strips_for_real_openai(self):
|
||||
body = await self.config.async_transform_request(
|
||||
model="gpt-4o",
|
||||
messages=self._cache_controlled_messages(),
|
||||
optional_params={},
|
||||
litellm_params={"custom_llm_provider": "openai", "api_base": None},
|
||||
headers={},
|
||||
)
|
||||
assert all("cache_control" not in m for m in body["messages"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue