mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(llm): preserve cache_control_injection_points at chat completion boundary
AmazonConverseConfig.transform_request reads cache_control_injection_points to append a cachePoint to the Bedrock tool list for location: tool_config. Stripping that key at the shared HTTP handler boundary silently disabled tool-config prompt caching on the bedrock/converse_like/ route, which is dispatched through base_llm_http_handler.completion. Split the strip set: chat-completion paths use a smaller set that preserves provider-consumed params, while the embedding boundary and splat-style transforms keep the full strip via strip_internal_params_from_request_body.
This commit is contained in:
parent
cb2badfebd
commit
8a14dc0dda
5 changed files with 94 additions and 3 deletions
|
|
@ -7,6 +7,7 @@ import httpx
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.internal_params import (
|
||||
LITELLM_CHAT_REQUEST_BODY_STRIP_PARAMS,
|
||||
LITELLM_INTERNAL_REQUEST_BODY_PARAMS,
|
||||
MCP_INTERNAL_PARAMS,
|
||||
)
|
||||
|
|
@ -481,6 +482,22 @@ def strip_internal_params_from_request_body(data: dict) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def strip_internal_params_from_chat_request_body(data: dict) -> dict:
|
||||
"""
|
||||
Strip variant for the chat-completion boundary that preserves keys consumed
|
||||
inside `transform_request` (currently `cache_control_injection_points`, which
|
||||
`AmazonConverseConfig` reads to append a `cachePoint` to Bedrock tool_config).
|
||||
Splat-style transforms still call `strip_internal_params_from_request_body`
|
||||
themselves before serialization, so the param never leaks into the wire body.
|
||||
"""
|
||||
if not isinstance(data, dict):
|
||||
return data
|
||||
|
||||
return {
|
||||
k: v for k, v in data.items() if k not in LITELLM_CHAT_REQUEST_BODY_STRIP_PARAMS
|
||||
}
|
||||
|
||||
|
||||
def redact_nested_match_and_regex_keys(
|
||||
payload: Union[dict, List[Any], str, None],
|
||||
) -> Union[dict, List[Any], str, None]:
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import litellm.litellm_core_utils
|
|||
import litellm.types
|
||||
import litellm.types.utils
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
strip_internal_params_from_request_body,
|
||||
strip_internal_params_from_chat_request_body,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig
|
||||
from litellm.llms.base_llm.image_variations.transformation import (
|
||||
|
|
@ -370,7 +370,9 @@ class BaseLLMAIOHTTPHandler:
|
|||
data = provider_config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=strip_internal_params_from_request_body(optional_params),
|
||||
optional_params=strip_internal_params_from_chat_request_body(
|
||||
optional_params
|
||||
),
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm._logging import _redact_string, verbose_logger
|
|||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
strip_internal_params_from_chat_request_body,
|
||||
strip_internal_params_from_request_body,
|
||||
)
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
|
@ -447,7 +448,9 @@ class BaseLLMHTTPHandler:
|
|||
data = provider_config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=strip_internal_params_from_request_body(optional_params),
|
||||
optional_params=strip_internal_params_from_chat_request_body(
|
||||
optional_params
|
||||
),
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -23,6 +23,17 @@ LITELLM_INTERNAL_REQUEST_BODY_PARAMS: frozenset[str] = frozenset(
|
|||
member.value for member in LiteLLMInternalParam
|
||||
)
|
||||
|
||||
LITELLM_CHAT_REQUEST_BODY_STRIP_PARAMS: frozenset[str] = (
|
||||
LITELLM_INTERNAL_REQUEST_BODY_PARAMS
|
||||
- frozenset({LiteLLMInternalParam.CACHE_CONTROL_INJECTION_POINTS.value})
|
||||
)
|
||||
"""Variant of `LITELLM_INTERNAL_REQUEST_BODY_PARAMS` for the chat-completion
|
||||
boundary. `cache_control_injection_points` is consumed inside `transform_request`
|
||||
by `AmazonConverseConfig` (it appends a `cachePoint` to the Bedrock tool list for
|
||||
``location: "tool_config"``), so it must reach the transform on the
|
||||
``converse_like/`` and other shared HTTP handler routes. Splat-style transforms
|
||||
that never consume it strip the full set themselves before serialization."""
|
||||
|
||||
MCP_INTERNAL_PARAMS: frozenset[str] = frozenset(
|
||||
{
|
||||
LiteLLMInternalParam.SKIP_MCP_HANDLER.value,
|
||||
|
|
|
|||
|
|
@ -79,6 +79,64 @@ def test_embedding_strips_internal_params_from_request_body():
|
|||
assert "output_dimension" in body
|
||||
|
||||
|
||||
def test_chat_boundary_preserves_cache_control_injection_points():
|
||||
"""Regression: the chat-completion boundary must NOT strip
|
||||
`cache_control_injection_points`. AmazonConverseConfig.transform_request
|
||||
consumes that key to append a `cachePoint` to Bedrock tool_config (used by
|
||||
the `converse_like/` route, which goes through this shared handler), so
|
||||
stripping it here silently disables tool-config prompt caching. Universal
|
||||
LiteLLM-internal knobs (skip_mcp_handler, fake_stream, ...) must still be
|
||||
stripped before the transform splats optional_params into the wire body."""
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.internal_params import LiteLLMInternalParam
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
provider_config = Mock()
|
||||
provider_config.should_fake_stream.return_value = False
|
||||
provider_config.validate_environment.return_value = {}
|
||||
provider_config.get_complete_url.return_value = "https://example.invalid/chat"
|
||||
provider_config.sign_request.return_value = ({}, None)
|
||||
|
||||
def _capture_transform_request(*, model, messages, optional_params, **_):
|
||||
captured["optional_params"] = optional_params
|
||||
raise RuntimeError("stop after capture")
|
||||
|
||||
provider_config.transform_request.side_effect = _capture_transform_request
|
||||
|
||||
seeded = {param.value: "internal" for param in LiteLLMInternalParam}
|
||||
seeded["cache_control_injection_points"] = [{"location": "tool_config"}]
|
||||
seeded["temperature"] = 0.5
|
||||
|
||||
with pytest.raises(RuntimeError, match="stop after capture"):
|
||||
handler.completion(
|
||||
model="anthropic.claude-3-5-haiku-20241022-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_base=None,
|
||||
custom_llm_provider="bedrock",
|
||||
model_response=ModelResponse(),
|
||||
encoding=None,
|
||||
logging_obj=Mock(),
|
||||
optional_params=seeded,
|
||||
timeout=10.0,
|
||||
litellm_params={},
|
||||
acompletion=False,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
|
||||
forwarded = captured["optional_params"]
|
||||
assert forwarded["cache_control_injection_points"] == [{"location": "tool_config"}]
|
||||
for param in LiteLLMInternalParam:
|
||||
if param is LiteLLMInternalParam.CACHE_CONTROL_INJECTION_POINTS:
|
||||
continue
|
||||
assert (
|
||||
param.value not in forwarded
|
||||
), f"{param.value} leaked past the chat-completion strip"
|
||||
|
||||
|
||||
def test_prepare_fake_stream_request():
|
||||
# Initialize the BaseLLMHTTPHandler
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue