mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
21a0b3c7d6
commit
c8720df299
2 changed files with 69 additions and 20 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue