fix(responses): avoid recursive cache-breakpoint stripping

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 01:54:18 +00:00
parent 21a0b3c7d6
commit c8720df299
2 changed files with 69 additions and 20 deletions

View file

@ -81,29 +81,36 @@ _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *Response
_CHAT_CONTENT_ITEM: Final = TypeAdapter(dict[str, object])
def _strip_prompt_cache_breakpoints_from_value(value: object) -> object:
if isinstance(value, dict):
content: Final = cast(dict[str, object], value) # cast-ok: isinstance narrows the recursive container
return {
key: _strip_prompt_cache_breakpoints_from_value(item)
for key, item in content.items()
if key != "prompt_cache_breakpoint"
}
def _strip_prompt_cache_breakpoint_from_content_block(value: object) -> object:
if not isinstance(value, dict):
return value
content_block: Final = cast(dict[str, object], value)
return {key: item for key, item in content_block.items() if key != "prompt_cache_breakpoint"}
def _strip_prompt_cache_breakpoints_from_content(value: object) -> object:
if isinstance(value, list):
return [
_strip_prompt_cache_breakpoints_from_value(item)
for item in cast(list[object], value) # cast-ok: isinstance narrows the recursive container
]
list_content: Final = cast(list[object], value)
return [_strip_prompt_cache_breakpoint_from_content_block(item) for item in list_content]
if isinstance(value, tuple):
return tuple(
_strip_prompt_cache_breakpoints_from_value(item)
for item in cast(tuple[object, ...], value) # cast-ok: isinstance narrows the recursive container
)
return value
tuple_content: Final = cast(tuple[object, ...], value)
return tuple(_strip_prompt_cache_breakpoint_from_content_block(item) for item in tuple_content)
return _strip_prompt_cache_breakpoint_from_content_block(value)
def _strip_prompt_cache_breakpoints_from_item(value: object) -> object:
if not isinstance(value, dict):
return value
input_item: Final = cast(dict[str, object], value)
return {
key: _strip_prompt_cache_breakpoints_from_content(item) if key in ("content", "output") else item
for key, item in input_item.items()
if key != "prompt_cache_breakpoint"
}
def _strip_prompt_cache_breakpoints(input_items: list[object]) -> list[object]:
return [_strip_prompt_cache_breakpoints_from_value(item) for item in input_items]
return [_strip_prompt_cache_breakpoints_from_item(item) for item in input_items]
def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]:

View file

@ -2,7 +2,7 @@ import datetime
import json
import os
import unittest
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, get_args
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast, get_args
from unittest.mock import ANY, MagicMock, Mock, patch
import httpx
@ -12,7 +12,7 @@ import litellm
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
from litellm.types.llms.openai import REASONING_EFFORT
from litellm.types.llms.openai import AllMessageValues, REASONING_EFFORT
if TYPE_CHECKING:
from openai.types.responses import ResponseOutputItem
@ -4291,6 +4291,48 @@ def test_prompt_cache_breakpoints_are_dropped_for_unsupported_models() -> None:
]
def test_prompt_cache_breakpoints_are_dropped_from_function_call_output_for_unsupported_models() -> None:
handler: Final = LiteLLMResponsesTransformationHandler()
cache_breakpoint: Final = {"mode": "explicit"}
messages: Final = cast(
list[AllMessageValues],
[
{
"role": "tool",
"tool_call_id": "call_1",
"content": [{"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint}],
}
],
)
request: Final = cast(
dict[str, object],
handler.transform_request(
model="gpt-5.4-mini",
messages=messages,
optional_params={},
litellm_params={},
headers={},
litellm_logging_obj=Mock(),
),
)
assert request["input"] == [
{
"type": "function_call_output",
"call_id": "call_1",
"output": [{"type": "input_text", "text": "Tool result"}],
}
]
assert messages == [
{
"role": "tool",
"tool_call_id": "call_1",
"content": [{"type": "text", "text": "Tool result", "prompt_cache_breakpoint": {"mode": "explicit"}}],
}
]
@pytest.mark.parametrize(
("litellm_params", "keep_marker"),
(({"base_model": "gpt-5.6"}, True), ({}, False)),