From c8720df299af84faddd99a9c6bad20a4ca927e96 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 01:54:18 +0000 Subject: [PATCH] fix(responses): avoid recursive cache-breakpoint stripping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../transformation.py | 43 +++++++++-------- ...responses_transformation_transformation.py | 46 ++++++++++++++++++- 2 files changed, 69 insertions(+), 20 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 58b3863b47b..96db88475ea 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -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]: diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 612232e1692..a45a56a5226 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -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)),