From ad6448d71b8f8f925a5d44122c3df3d06f3511af Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 00:32:30 +0000 Subject: [PATCH 1/9] fix(responses): keep prompt_cache_breakpoint markers in the chat to responses bridge Preserve cache-breakpoint markers through the bridge for supported models and drop them for models without breakpoint support Co-authored-by: Simon Sorg Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../transformation.py | 79 +- ...esponses_bridge_prompt_cache_breakpoint.py | 911 ++++++++++++++++++ ...responses_transformation_transformation.py | 106 ++ 3 files changed, 1079 insertions(+), 17 deletions(-) create mode 100644 tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 391dbc44eec..649f98b5d9e 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -19,13 +19,15 @@ from openai.types.responses.response_input_param import ( from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam from openai.types.responses.tool_param import FunctionToolParam -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter import litellm from litellm import ModelResponse from litellm._logging import verbose_logger +from litellm.integrations.anthropic_cache_control_hook import supports_openai_prompt_cache_breakpoint from litellm.litellm_core_utils.prompt_templates.common_utils import ( responses_reasoning_items_from_thinking_blocks, + with_prompt_cache_breakpoint, ) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.bridges.completion_transformation import ( @@ -76,6 +78,32 @@ _CHAT_COMPLETION_FIELDS: Final = frozenset((*ModelResponse.model_fields, "usage" _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *ResponsesAPIResponse.model_fields)) - frozenset( ChatCompletion.model_fields ) +_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" + } + 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 + ] + 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 + + +def _strip_prompt_cache_breakpoints(input_items: list[object]) -> list[object]: + return [_strip_prompt_cache_breakpoints_from_value(item) for item in input_items] def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]: @@ -588,24 +616,32 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): litellm_logging_obj: "LiteLLMLoggingObj", client: object | None = None, ) -> dict: - ( - input_items, - instructions, - ) = self.convert_chat_completion_messages_to_responses_api(messages) - + converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(messages) + supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) + input_items_without_unsupported_markers: Final = ( + converted_input_items + if supports_prompt_cache_breakpoint + else _strip_prompt_cache_breakpoints(converted_input_items) + ) # OpenAI's Responses API rejects an empty input. For a system-only # request, carry the system message as a system-role input item instead # of instructions, mirroring how non-string system content is already # handled in convert_chat_completion_messages_to_responses_api. - if not input_items and instructions is not None: - input_items = [ + is_system_only_request: Final = ( + not input_items_without_unsupported_markers and converted_instructions is not None + ) + input_items: Final = ( + [ { "type": "message", "role": "system", - "content": [{"type": "input_text", "text": instructions}], + "content": [{"type": "input_text", "text": converted_instructions}], } ] - instructions = None + if is_system_only_request + else input_items_without_unsupported_markers + ) + instructions: Final = None if is_system_only_request else converted_instructions optional_params = self._extract_extra_body_params(optional_params) @@ -1060,16 +1096,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Handle multimodal content original_type = item.get("type") if original_type == "text": - converted = self._convert_content_str_to_input_text(item.get("text", ""), role) + converted = with_prompt_cache_breakpoint( + self._convert_content_str_to_input_text(item.get("text", ""), role), + _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), + ) result.append(converted) verbose_logger.debug("Chat provider: text -> %s", converted) elif original_type == "image_url": # Map to responses API image format - converted = cast( - dict, - self._convert_content_to_responses_format_image( - cast(ChatCompletionImageObject, item), role + converted = with_prompt_cache_breakpoint( + dict( + self._convert_content_to_responses_format_image( + cast(ChatCompletionImageObject, item), # cast-ok: image_url type tag was checked + role, + ) ), + _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) @@ -1081,8 +1123,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): result.append(converted) verbose_logger.debug("Chat provider: image -> %s", converted) elif item_type == "file": - converted = _input_file_from_file_value( - cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked + converted = with_prompt_cache_breakpoint( + _input_file_from_file_value( + cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked + ), + _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: file -> %s", converted) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py new file mode 100644 index 00000000000..75d8ffac9f7 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py @@ -0,0 +1,911 @@ +from __future__ import annotations + +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Iterable, Mapping +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final, Literal, TypeAlias, cast +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolUnionParam +from pydantic import JsonValue, TypeAdapter + +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from litellm.responses.utils import ResponsesAPIRequestUtils as _RU + +_MODEL: Final = "openai/gpt-5.6" +_UNSUPPORTED_MODEL: Final = "openai/gpt-5.4-mini" +_CONFIG_MODEL: Final = "responses-bridge-cache-breakpoint-chaos" +_API_KEY: Final = "synthetic-responses-bridge-key" +_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})") +_STARTED_WORKER: Final[re.Pattern[str]] = re.compile(r"Started server process \[(\d+)\]") +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_BREAKPOINT: Final[dict[str, JsonValue]] = {"mode": "explicit"} +_TOOLS: Final[list[JsonValue]] = [ + { + "type": "function", + "function": { + "name": "synthetic_tool", + "description": "Synthetic bridge test tool", + "parameters": {"type": "object", "properties": {}}, + }, + } +] +_IMAGE_URL: Final = "data:image/png;base64,aGVsbG8=" +_ClientKind: TypeAlias = Literal["openai_sync", "openai_async", "httpx"] +_Surface: TypeAlias = Literal["chat", "responses"] + + +@dataclass(frozen=True, slots=True) +class _Call: + surface: _Surface + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + response_id: str | None + text: str + + +def _response_id(marker: str) -> str: + return f"resp_{marker}" + + +def _request_marker(request: Request) -> str: + match: Final = _MARKER.search(request.body) + assert match is not None, request.body + return match.group(1).decode() + + +def _contains_breakpoint(value: JsonValue) -> bool: + if isinstance(value, dict): + return "prompt_cache_breakpoint" in value or any(_contains_breakpoint(item) for item in value.values()) + if isinstance(value, list): + return any(_contains_breakpoint(item) for item in value) + return False + + +def _responses_body(marker: str) -> dict[str, JsonValue]: + response_id: Final = _response_id(marker) + return _JSON_OBJECT.validate_python( + { + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.6", + "output": [ + { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": f"answer marker-{marker}", "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}, + } + ) + + +def _responses_reply(request: Request, *, reject_breakpoints: bool = False) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return Reply(body=b'{"object":"list","data":[{"id":"gpt-5.6","object":"model"}]}') + body: Final = _JSON_OBJECT.validate_json(request.body) + if reject_breakpoints and _contains_breakpoint(body): + return Reply( + status=400, + body=json.dumps( + { + "error": { + "message": "prompt_cache_breakpoint is not supported on this model", + "type": "invalid_request_error", + "param": None, + "code": None, + } + } + ).encode(), + ) + marker: Final = _request_marker(request) + stream: Final = body.get("stream") is True + response: Final = _responses_body(marker) + if not stream: + return Reply(body=json.dumps(response).encode()) + created: Final = { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + } + delta: Final = { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": f"msg_{marker}", + "output_index": 0, + "content_index": 0, + "delta": f"answer marker-{marker}", + } + completed: Final = {"type": "response.completed", "sequence_number": 2, "response": response} + events: Final = (created, delta, completed) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _prompt(marker: str, label: str) -> str: + return f"{label} marker-{marker}" + + +def _simple_chat_body( + model: str, + marker: str, + *, + stream: bool = False, + marked: bool = True, + system_as_string: bool = False, + prompt_cache_options: dict[str, JsonValue] | None = None, +) -> dict[str, JsonValue]: + marker_field: Final = {"prompt_cache_breakpoint": _BREAKPOINT} if marked else {} + user: Final = [{"type": "text", "text": _prompt(marker, "user"), **marker_field}] + messages: Final = ( + [{"role": "system", "content": _prompt(marker, "system")}, {"role": "user", "content": user}] + if system_as_string + else [{"role": "user", "content": user}] + ) + return _JSON_OBJECT.validate_python( + { + "model": model, + "messages": messages, + "tools": _TOOLS, + "reasoning_effort": "low", + "stream": stream, + "num_retries": 0, + **({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}), + } + ) + + +def _multimodal_chat_body(model: str, marker: str, stream: bool) -> dict[str, JsonValue]: + return _JSON_OBJECT.validate_python( + { + "model": model, + "messages": [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": _prompt(marker, "system"), + "prompt_cache_breakpoint": _BREAKPOINT, + } + ], + }, + { + "role": "user", + "content": [ + {"type": "text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, + { + "type": "image_url", + "image_url": {"url": _IMAGE_URL}, + "prompt_cache_breakpoint": _BREAKPOINT, + }, + { + "type": "file", + "file": {"file_id": "file-abc"}, + "prompt_cache_breakpoint": _BREAKPOINT, + }, + {"type": "text", "text": "unmarked extra text"}, + ], + }, + ], + "tools": _TOOLS, + "reasoning_effort": "low", + "stream": stream, + "num_retries": 0, + "prompt_cache_options": {"mode": "explicit"}, + } + ) + + +def _expected_multimodal_input(marker: str) -> list[JsonValue]: + return [ + { + "type": "message", + "role": "system", + "content": [ + {"type": "input_text", "text": _prompt(marker, "system"), "prompt_cache_breakpoint": _BREAKPOINT} + ], + }, + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, + { + "type": "input_image", + "image_url": _IMAGE_URL, + "detail": "auto", + "prompt_cache_breakpoint": _BREAKPOINT, + }, + {"type": "input_file", "file_id": "file-abc", "prompt_cache_breakpoint": _BREAKPOINT}, + {"type": "input_text", "text": "unmarked extra text"}, + ], + }, + ] + + +def _simple_expected_input(marker: str, *, marked: bool) -> list[JsonValue]: + text_block: Final = {"type": "input_text", "text": _prompt(marker, "user")} + return [ + { + "type": "message", + "role": "user", + "content": [{**text_block, **({"prompt_cache_breakpoint": _BREAKPOINT} if marked else {})}], + } + ] + + +def _request_body(request: Request) -> dict[str, JsonValue]: + assert request.method == "POST" and request.target == "/v1/responses", request.target + return _JSON_OBJECT.validate_json(request.body) + + +def _decoded_response_id(response_id: str) -> str: + decoded: Final = _RU._decode_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # reuse ID decoder + response_id + ) + raw_response_id: Final = decoded.get("response_id") + assert isinstance(raw_response_id, str), decoded + return raw_response_id + + +def _spend_request_id_matches( + row: Mapping[str, JsonValue], + caller_response_id: str, + peer_response_id: str, + surface: _Surface, +) -> bool: + request_id: Final = row.get("request_id") + if not isinstance(request_id, str): + return False + match surface: + case "responses": + return request_id == caller_response_id + case "chat": + return _decoded_response_id(request_id) == peer_response_id + + +def _spend_rows( + model: str, + caller_response_id: str, + peer_response_id: str, + surface: _Surface, +) -> tuple[dict[str, JsonValue], ...]: + def matching_rows(rows: list[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + row for row in rows if _spend_request_id_matches(row, caller_response_id, peer_response_id, surface) + ) + + rows: Final = eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda candidates: len(matching_rows(candidates)) == 1, + seconds=60, + ) + matched: Final = matching_rows(rows) + assert len(matched) == 1, matched + return matched + + +def _response_id_from_chat_stream(text: str) -> str: + payloads: Final = tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: {") + ) + assert payloads, text + response_id: Final = payloads[0].get("id") + assert isinstance(response_id, str), payloads[0] + return response_id + + +def _extra_body(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return { + key: value + for key, value in body.items() + if key not in {"model", "messages", "tools", "reasoning_effort", "stream", "num_retries"} + } + + +def _sync_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: + base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" + model: Final = str(body["model"]) + messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) + tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) + extras: Final = _extra_body(body) + with OpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: + if stream: + response_stream: Final = client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=True, + extra_body=extras, + ) + chunks: Final = tuple(response_stream) + assert chunks + return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") + response: Final = client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=False, + extra_body=extras, + ) + return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") + + +async def _async_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: + base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" + model: Final = str(body["model"]) + messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) + tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) + extras: Final = _extra_body(body) + async with AsyncOpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: + if stream: + response_stream: Final = await client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=True, + extra_body=extras, + ) + chunks: Final = tuple([chunk async for chunk in response_stream]) + assert chunks + return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") + response: Final = await client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=False, + extra_body=extras, + ) + return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") + + +async def _serve_chat( + gateway: Gateway, + body: dict[str, JsonValue], + client_kind: _ClientKind, + stream: bool, +) -> _Served: + match client_kind: + case "openai_sync": + return _sync_sdk_chat(gateway, body, stream) + case "openai_async": + return await _async_sdk_chat(gateway, body, stream) + case "httpx": + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + return await _raw_call( + client, + "/v1/chat/completions", + body, + _Call("chat", stream, _request_marker_from_body(body)), + ) + + +def _request_marker_from_body(body: Mapping[str, JsonValue]) -> str: + match: Final = _MARKER.search(json.dumps(body).encode()) + assert match is not None, body + return match.group(1).decode() + + +async def _raw_call( + client: httpx.AsyncClient, + path: str, + body: Mapping[str, JsonValue], + call: _Call, +) -> _Served: + async with client.stream( + "POST", + path, + json=body, + headers={"Authorization": f"Bearer {client.headers['Authorization'].removeprefix('Bearer ')}"}, + ) as response: + content: Final = await response.aread() + status: Final = response.status_code + text: Final = content.decode() + response_id: Final = ( + _response_id_from_chat_stream(text) + if status == 200 and call.surface == "chat" and call.stream + else _JSON_OBJECT.validate_json(content).get("id") + if status == 200 + else None + ) + return _Served(call, status, response_id if isinstance(response_id, str) else None, text) + + +async def _send_call( + client: httpx.AsyncClient, + model: str, + call: _Call, +) -> _Served: + body: Final = ( + _simple_chat_body(model, call.marker, stream=call.stream) + if call.surface == "chat" + else { + "model": model, + "input": _simple_expected_input(call.marker, marked=True), + "stream": call.stream, + "num_retries": 0, + } + ) + path: Final = "/v1/chat/completions" if call.surface == "chat" else "/v1/responses" + try: + return await _raw_call(client, path, body, call) + except httpx.TransportError as error: + return _Served(call, 0, None, f"{type(error).__name__}: {error}") + + +async def _burst( + base_url: str, + key: str, + model: str, + calls: tuple[_Call, ...], +) -> tuple[_Served, ...]: + async with httpx.AsyncClient( + base_url=base_url, + headers={"Authorization": f"Bearer {key}"}, + timeout=20, + trust_env=False, + limits=httpx.Limits(max_connections=100), + ) as client: + return tuple(await asyncio.gather(*(_send_call(client, model, call) for call in calls))) + + +def _calls(count: int) -> tuple[_Call, ...]: + surfaces: Final[tuple[_Surface, ...]] = ("chat", "chat", "responses") + return tuple( + _Call( + surface=surfaces[index % len(surfaces)], + stream=index % 3 == 1, + marker=uuid.uuid4().hex, + ) + for index in range(count) + ) + + +def _requests_for_marker(requests: tuple[Request, ...], marker: str) -> tuple[Request, ...]: + return tuple( + request + for request in requests + if request.method == "POST" and request.target == "/v1/responses" and _request_marker(request) == marker + ) + + +def _peer_request_has_marker(request: Request, marker: str) -> bool: + body: Final = _request_body(request) + return _contains_breakpoint(body) and _request_marker(request) == marker + + +def _peer_marker_matches_response(served: _Served, requests: tuple[Request, ...]) -> bool: + peer_requests: Final = _requests_for_marker(requests, served.call.marker) + assert len(peer_requests) == 1, (served, peer_requests) + (peer_request,) = peer_requests + return _peer_request_has_marker(peer_request, served.call.marker) + + +def _assert_spend_for_result(served: _Served, model: str) -> None: + assert served.status == 200 and served.response_id is not None, served + peer_response_id: Final = _response_id(served.call.marker) + match served.call.surface: + case "responses": + (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) + case "chat": + assert _decoded_response_id(served.response_id) == peer_response_id, served + (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) + request_id: Final = row.get("request_id") + assert isinstance(request_id, str), row + match served.call.surface: + case "responses": + assert request_id == served.response_id, row + case "chat": + assert _decoded_response_id(request_id) == served.response_id, row + + +@pytest.mark.parametrize("client_kind", ("openai_sync", "openai_async", "httpx")) +@pytest.mark.parametrize("stream", (False, True)) +async def test_caller_prompt_cache_breakpoints_survive_chat_to_responses_bridge( + gateway: Gateway, + client_kind: _ClientKind, + stream: bool, +) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_MODEL, api_base=wire.url + "/v1") + body: Final = _multimodal_chat_body(model, marker, stream) + served: Final = await _serve_chat(gateway, body, client_kind, stream) + assert served.status == 200, served.text + _assert_spend_for_result(served, model) + (peer_request,) = wire.drain() + peer_body: Final = _request_body(peer_request) + assert peer_body["input"] == _expected_multimodal_input(marker), peer_body + assert peer_body["prompt_cache_options"] == {"mode": "explicit"}, peer_body + + +def _expected_uninjected_system_bridge_body( + marker: str, + prompt_cache_options: dict[str, JsonValue] | None = None, +) -> dict[str, JsonValue]: + return { + "input": _simple_expected_input(marker, marked=False), + "instructions": _prompt(marker, "system"), + "model": "gpt-5.6", + "reasoning": {"effort": "low"}, + "stream": False, + "tools": [ + { + "type": "function", + "name": "synthetic_tool", + "parameters": {"type": "object", "properties": {}}, + "strict": None, + "description": "Synthetic bridge test tool", + } + ], + **({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}), + } + + +async def test_deployment_cache_control_injection_without_options_is_unchanged(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=_MODEL, + api_base=wire.url + "/v1", + cache_control_injection_points=[{"location": "message", "role": "system"}], + ) + body: Final = _simple_chat_body(model, marker, system_as_string=True, marked=False) + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + served: Final = await _raw_call(client, "/v1/chat/completions", body, _Call("chat", False, marker)) + assert served.status == 200, served.text + _assert_spend_for_result(served, model) + (peer_request,) = wire.drain() + peer_body: Final = _request_body(peer_request) + assert peer_body == _expected_uninjected_system_bridge_body(marker), peer_body + assert not _contains_breakpoint(peer_body), peer_body + assert "prompt_cache_options" not in peer_body, peer_body + + +async def test_deployment_prompt_cache_options_override_is_unchanged(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + options: Final[dict[str, JsonValue]] = {"mode": "implicit", "ttl": "30m"} + with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=_MODEL, + api_base=wire.url + "/v1", + prompt_cache_options=options, + ) + body: Final = _simple_chat_body(model, marker, system_as_string=True, marked=False) + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + served: Final = await _raw_call(client, "/v1/chat/completions", body, _Call("chat", False, marker)) + assert served.status == 200, served.text + _assert_spend_for_result(served, model) + (peer_request,) = wire.drain() + peer_body: Final = _request_body(peer_request) + assert peer_body == _expected_uninjected_system_bridge_body(marker, options), peer_body + assert not _contains_breakpoint(peer_body), peer_body + + +async def test_unmarked_bridge_and_direct_responses_marker_are_forwarded_unchanged(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_MODEL, api_base=wire.url + "/v1") + unmarked_body: Final = _simple_chat_body(model, marker, marked=False) + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + unmarked: Final = await _raw_call( + client, + "/v1/chat/completions", + unmarked_body, + _Call("chat", False, marker), + ) + assert unmarked.status == 200, unmarked.text + _assert_spend_for_result(unmarked, model) + (unmarked_peer,) = wire.drain() + unmarked_body_at_peer: Final = _request_body(unmarked_peer) + assert not _contains_breakpoint(unmarked_body_at_peer), unmarked_body_at_peer + assert "prompt_cache_options" not in unmarked_body_at_peer, unmarked_body_at_peer + + direct_marker: Final = uuid.uuid4().hex + direct_input: Final = [ + { + "type": "message", + "role": "user", + "content": [ + { + "type": "input_text", + "text": _prompt(direct_marker, "direct"), + "prompt_cache_breakpoint": _BREAKPOINT, + } + ], + } + ] + direct_body: Final = _JSON_OBJECT.validate_python({"model": model, "input": direct_input, "store": False}) + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + direct: Final = await _raw_call( + client, + "/v1/responses", + direct_body, + _Call("responses", False, direct_marker), + ) + assert direct.status == 200, direct.text + _assert_spend_for_result(direct, model) + (direct_peer,) = wire.drain() + assert _request_body(direct_peer)["input"] == direct_input, _request_body(direct_peer) + + +@pytest.mark.parametrize("stream", (False, True)) +async def test_unsupported_model_drops_breakpoints_without_rejecting_the_request( + gateway: Gateway, + stream: bool, +) -> None: + marker: Final = uuid.uuid4().hex + with ( + wire_server(lambda request: _responses_reply(request, reject_breakpoints=True)) as wire, + gateway.scenario() as scenario, + ): + model: Final = scenario.model(model=_UNSUPPORTED_MODEL, api_base=wire.url + "/v1") + body: Final = _simple_chat_body(model, marker, stream=stream) + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + served: Final = await _raw_call(client, "/v1/chat/completions", body, _Call("chat", stream, marker)) + assert served.status == 200, served.text + _assert_spend_for_result(served, model) + (peer_request,) = wire.drain() + peer_body: Final = _request_body(peer_request) + assert peer_body["input"] == _simple_expected_input(marker, marked=False), peer_body + assert not _contains_breakpoint(peer_body), peer_body + + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + base_config: Final = _JSON_OBJECT.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + config: Final = { + **base_config, + "model_list": [ + { + "model_name": _CONFIG_MODEL, + "litellm_params": { + "model": _MODEL, + "api_base": wire.url + "/v1", + "api_key": _API_KEY, + }, + }, + ], + } + path: Final = tmp_path / "responses-bridge-cache-breakpoint-chaos.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _open_upstream_connections(pid: int, port: int) -> int: + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(180) +async def test_worker_and_peer_outages_preserve_markers_and_recover( + gateway: Gateway, + tmp_path: Path, +) -> None: + calls: Final = _calls(30) + release: Final = threading.Event() + early_release: Final = threading.Event() + outage_release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + early_calls: Final = calls[:10] + early_markers: Final = frozenset(call.marker for call in early_calls) + + def held(request: Request) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return _responses_reply(request) + marker: Final = _request_marker(request) + held_markers.put(marker) + gate: Final = early_release if marker in early_markers else release + assert gate.wait(timeout=60), "The worker-kill burst was never released" + return _responses_reply(request) + + with ExitStack() as peer_stack: + wire: Final = peer_stack.enter_context(wire_server(held)) + config: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + try: + candidate: Final = owned.gateway + workers: Final[tuple[int, ...]] = eventually( + lambda: tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + async with httpx.AsyncClient( + base_url=str(candidate.client.base_url), + headers={"Authorization": f"Bearer {candidate.key}"}, + timeout=20, + trust_env=False, + limits=httpx.Limits(max_connections=100), + ) as client: + burst_tasks: Final = tuple( + asyncio.create_task(_send_call(client, _CONFIG_MODEL, call)) for call in calls + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == len(calls), 60) + early_release.set() + early_served: Final = await asyncio.gather(*burst_tasks[: len(early_calls)]) + early_successful: Final = tuple(item for item in early_served if item.status == 200) + for item in early_successful: + _assert_spend_for_result(item, _CONFIG_MODEL) + upstream_port_value: Final = urlsplit(wire.url).port + assert upstream_port_value is not None + upstream_port: Final = upstream_port_value + active_by_worker: Final = eventually( + lambda: {pid: _open_upstream_connections(pid, upstream_port) for pid in workers}, + lambda counts: sum(counts.values()) == len(calls) - len(early_calls), + seconds=30, + ) + victim_pid: Final = max(workers, key=active_by_worker.__getitem__) + survivor_pids: Final = tuple(pid for pid in workers if pid != victim_pid) + assert active_by_worker[victim_pid] > 0 and len(survivor_pids) == 1, active_by_worker + (survivor_pid,) = survivor_pids + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + remaining_served: Final = await asyncio.gather(*burst_tasks[len(early_calls) :]) + served: Final = (*early_served, *remaining_served) + successful: Final = tuple(item for item in served if item.status == 200) + connection_errors: Final = tuple(item for item in served if item.status == 0) + print(f"worker-kill burst: {len(successful)} HTTP 200, {len(connection_errors)} connection errors") + assert len(successful) + len(connection_errors) == len(calls), { + "successes": len(successful), + "connection_errors": len(connection_errors), + "responses": served, + } + assert successful and connection_errors, { + "successes": len(successful), + "connection_errors": len(connection_errors), + } + follow_ups: Final = ( + _Call("chat", False, uuid.uuid4().hex), + _Call("responses", False, uuid.uuid4().hex), + ) + recovered: Final = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + follow_ups, + ) + assert all(item.status == 200 for item in recovered), recovered + assert psutil.pid_exists(survivor_pid), survivor_pid + received_after_worker_kill: Final = wire.drain() + worker_marker_failures: Final = tuple( + item.call.marker + for item in (*successful, *recovered) + if not _peer_marker_matches_response(item, received_after_worker_kill) + ) + for item in (*successful, *recovered): + assert item.response_id is not None, item + _assert_spend_for_result(item, _CONFIG_MODEL) + + peer_stack.close() + outage_seen: Final[SimpleQueue[str]] = SimpleQueue() + + def outage(request: Request) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return _responses_reply(request) + outage_seen.put(_request_marker(request)) + assert outage_release.wait(timeout=60), "The peer-outage burst was never stopped" + return Reply( + status=503, + body=b'{"error":{"message":"synthetic peer outage","type":"server_error"}}', + ) + + peer_stack.enter_context(wire_server(outage, port=upstream_port)) + outage_calls: Final = _calls(12) + outage_burst: Final = asyncio.create_task( + _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, outage_calls) + ) + await asyncio.to_thread(eventually, outage_seen.qsize, lambda size: size == len(outage_calls), 30) + outage_release.set() + peer_stack.close() + outage_served: Final = await outage_burst + assert len(outage_served) == len(outage_calls), outage_served + assert all(item.status >= 400 and "error" in item.text.lower() for item in outage_served), outage_served + down_call: Final = _Call("chat", False, uuid.uuid4().hex) + (down_response,) = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + (down_call,), + ) + assert down_response.status >= 400 and "error" in down_response.text.lower(), down_response + + restarted_wire: Final = peer_stack.enter_context(wire_server(_responses_reply, port=upstream_port)) + recovery_calls: Final = ( + _Call("chat", False, uuid.uuid4().hex), + _Call("responses", False, uuid.uuid4().hex), + ) + recovered_after_peer_restart: Final = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + recovery_calls, + ) + assert all(item.status == 200 for item in recovered_after_peer_restart), recovered_after_peer_restart + restarted_requests: Final = restarted_wire.drain() + recovery_marker_failures: Final = tuple( + item.call.marker + for item in recovered_after_peer_restart + if not _peer_marker_matches_response(item, restarted_requests) + ) + for item in recovered_after_peer_restart: + assert item.response_id is not None, item + _assert_spend_for_result(item, _CONFIG_MODEL) + assert not (*worker_marker_failures, *recovery_marker_failures), { + "worker_marker_failures": worker_marker_failures, + "recovery_marker_failures": recovery_marker_failures, + } + finally: + release.set() + outage_release.set() 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 282b84104a6..c6a4cdaed39 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 @@ -4185,6 +4185,112 @@ def _system_input_item(text: str) -> dict[str, object]: return {"type": "message", "role": "system", "content": [{"type": "input_text", "text": text}]} +@pytest.mark.parametrize( + ("content_block", "expected_content"), + [ + ( + {"type": "text", "text": "Stable prefix"}, + {"type": "input_text", "text": "Stable prefix"}, + ), + ( + {"type": "image_url", "image_url": "https://example.com/image.png"}, + {"type": "input_image", "image_url": "https://example.com/image.png", "detail": "auto"}, + ), + ( + {"type": "file", "file": {"file_id": "file-123"}}, + {"type": "input_file", "file_id": "file-123"}, + ), + ], + ids=("text", "image_url", "file"), +) +def test_prompt_cache_breakpoint_survives_chat_to_responses_conversion( + content_block: dict[str, object], expected_content: dict[str, object] +) -> None: + handler: Final = LiteLLMResponsesTransformationHandler() + cache_breakpoint: Final = {"mode": "explicit"} + marked_content: Final = {**content_block, "prompt_cache_breakpoint": cache_breakpoint} + + request: Final = handler.transform_request( + model="gpt-5.6", + messages=[ + { + "role": "user", + "content": [marked_content], + } + ], + optional_params={"prompt_cache_options": cache_breakpoint}, + litellm_params={}, + headers={}, + litellm_logging_obj=Mock(), + ) + + assert request["input"][0] == { + "type": "message", + "role": "user", + "content": [{**expected_content, "prompt_cache_breakpoint": cache_breakpoint}], + } + assert request["prompt_cache_options"] == cache_breakpoint + + +def test_prompt_cache_breakpoints_are_dropped_for_unsupported_models() -> None: + handler: Final = LiteLLMResponsesTransformationHandler() + cache_breakpoint: Final = {"mode": "explicit"} + marked_content: Final = [ + {"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": cache_breakpoint}, + { + "type": "image_url", + "image_url": "https://example.com/image.png", + "prompt_cache_breakpoint": cache_breakpoint, + }, + { + "type": "file", + "file": {"file_id": "file-123"}, + "prompt_cache_breakpoint": cache_breakpoint, + }, + ] + messages: Final = [{"role": "user", "content": marked_content}] + + request: Final = handler.transform_request( + model="gpt-5.4-mini", + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + litellm_logging_obj=Mock(), + ) + + assert request["input"] == [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "Stable prefix"}, + {"type": "input_image", "image_url": "https://example.com/image.png", "detail": "auto"}, + {"type": "input_file", "file_id": "file-123"}, + ], + } + ] + assert "prompt_cache_options" not in request + assert messages == [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": {"mode": "explicit"}}, + { + "type": "image_url", + "image_url": "https://example.com/image.png", + "prompt_cache_breakpoint": {"mode": "explicit"}, + }, + { + "type": "file", + "file": {"file_id": "file-123"}, + "prompt_cache_breakpoint": {"mode": "explicit"}, + }, + ], + } + ] + + def test_mid_conversation_system_string_stays_in_input_after_a_user_turn(): handler: Final = LiteLLMResponsesTransformationHandler() From 21a0b3c7d66d9a932cd511277e5399b25ffc3cb1 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 01:22:57 +0000 Subject: [PATCH 2/9] fix(responses): honor base_model when gating bridge cache breakpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../transformation.py | 5 +- ...esponses_bridge_prompt_cache_breakpoint.py | 505 ++++++++++++ ...esponses_bridge_prompt_cache_breakpoint.py | 756 +----------------- ...es_bridge_prompt_cache_breakpoint_chaos.py | 229 ++++++ ...responses_transformation_transformation.py | 36 + 5 files changed, 797 insertions(+), 734 deletions(-) create mode 100644 tests/integration/providers/_responses_bridge_prompt_cache_breakpoint.py create mode 100644 tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_chaos.py diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 649f98b5d9e..58b3863b47b 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -617,7 +617,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): client: object | None = None, ) -> dict: converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(messages) - supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) + base_model: Final = litellm_params.get("base_model") + supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) or ( + isinstance(base_model, str) and bool(base_model) and supports_openai_prompt_cache_breakpoint(base_model) + ) input_items_without_unsupported_markers: Final = ( converted_input_items if supports_prompt_cache_breakpoint diff --git a/tests/integration/providers/_responses_bridge_prompt_cache_breakpoint.py b/tests/integration/providers/_responses_bridge_prompt_cache_breakpoint.py new file mode 100644 index 00000000000..bc386965254 --- /dev/null +++ b/tests/integration/providers/_responses_bridge_prompt_cache_breakpoint.py @@ -0,0 +1,505 @@ +from __future__ import annotations + +import asyncio +import json +import re +import uuid +from collections.abc import Iterable, Mapping +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias, cast + +import httpx +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolUnionParam +from pydantic import JsonValue, TypeAdapter + +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request +from litellm.responses.utils import ResponsesAPIRequestUtils as _RU + +_MODEL: Final = "openai/gpt-5.6" + +_UNSUPPORTED_MODEL: Final = "openai/gpt-5.4-mini" + +_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})") + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + +_BREAKPOINT: Final[dict[str, JsonValue]] = {"mode": "explicit"} + +_TOOLS: Final[list[JsonValue]] = [ + { + "type": "function", + "function": { + "name": "synthetic_tool", + "description": "Synthetic bridge test tool", + "parameters": {"type": "object", "properties": {}}, + }, + } +] + +_IMAGE_URL: Final = "data:image/png;base64,aGVsbG8=" + +_ClientKind: TypeAlias = Literal["openai_sync", "openai_async", "httpx"] + +_Surface: TypeAlias = Literal["chat", "responses"] + +@dataclass(frozen=True, slots=True) +class _Call: + surface: _Surface + stream: bool + marker: str + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + response_id: str | None + text: str + +def _response_id(marker: str) -> str: + return f"resp_{marker}" + +def _request_marker(request: Request) -> str: + match: Final = _MARKER.search(request.body) + assert match is not None, request.body + return match.group(1).decode() + +def _contains_breakpoint(value: JsonValue) -> bool: + if isinstance(value, dict): + return "prompt_cache_breakpoint" in value or any(_contains_breakpoint(item) for item in value.values()) + if isinstance(value, list): + return any(_contains_breakpoint(item) for item in value) + return False + +def _responses_body(marker: str) -> dict[str, JsonValue]: + response_id: Final = _response_id(marker) + return _JSON_OBJECT.validate_python( + { + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.6", + "output": [ + { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": f"answer marker-{marker}", "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}, + } + ) + +def _responses_reply(request: Request, *, reject_breakpoints: bool = False) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return Reply(body=b'{"object":"list","data":[{"id":"gpt-5.6","object":"model"}]}') + body: Final = _JSON_OBJECT.validate_json(request.body) + if reject_breakpoints and _contains_breakpoint(body): + return Reply( + status=400, + body=json.dumps( + { + "error": { + "message": "prompt_cache_breakpoint is not supported on this model", + "type": "invalid_request_error", + "param": None, + "code": None, + } + } + ).encode(), + ) + marker: Final = _request_marker(request) + stream: Final = body.get("stream") is True + response: Final = _responses_body(marker) + if not stream: + return Reply(body=json.dumps(response).encode()) + created: Final = { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + } + delta: Final = { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": f"msg_{marker}", + "output_index": 0, + "content_index": 0, + "delta": f"answer marker-{marker}", + } + completed: Final = {"type": "response.completed", "sequence_number": 2, "response": response} + events: Final = (created, delta, completed) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + +def _prompt(marker: str, label: str) -> str: + return f"{label} marker-{marker}" + +def _simple_chat_body( + model: str, + marker: str, + *, + stream: bool = False, + marked: bool = True, + system_as_string: bool = False, + prompt_cache_options: dict[str, JsonValue] | None = None, +) -> dict[str, JsonValue]: + marker_field: Final = {"prompt_cache_breakpoint": _BREAKPOINT} if marked else {} + user: Final = [{"type": "text", "text": _prompt(marker, "user"), **marker_field}] + messages: Final = ( + [{"role": "system", "content": _prompt(marker, "system")}, {"role": "user", "content": user}] + if system_as_string + else [{"role": "user", "content": user}] + ) + return _JSON_OBJECT.validate_python( + { + "model": model, + "messages": messages, + "tools": _TOOLS, + "reasoning_effort": "low", + "stream": stream, + "num_retries": 0, + **({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}), + } + ) + +def _multimodal_chat_body(model: str, marker: str, stream: bool) -> dict[str, JsonValue]: + return _JSON_OBJECT.validate_python( + { + "model": model, + "messages": [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": _prompt(marker, "system"), + "prompt_cache_breakpoint": _BREAKPOINT, + } + ], + }, + { + "role": "user", + "content": [ + {"type": "text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, + { + "type": "image_url", + "image_url": {"url": _IMAGE_URL}, + "prompt_cache_breakpoint": _BREAKPOINT, + }, + { + "type": "file", + "file": {"file_id": "file-abc"}, + "prompt_cache_breakpoint": _BREAKPOINT, + }, + {"type": "text", "text": "unmarked extra text"}, + ], + }, + ], + "tools": _TOOLS, + "reasoning_effort": "low", + "stream": stream, + "num_retries": 0, + "prompt_cache_options": {"mode": "explicit"}, + } + ) + +def _expected_multimodal_input(marker: str) -> list[JsonValue]: + return [ + { + "type": "message", + "role": "system", + "content": [ + {"type": "input_text", "text": _prompt(marker, "system"), "prompt_cache_breakpoint": _BREAKPOINT} + ], + }, + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, + { + "type": "input_image", + "image_url": _IMAGE_URL, + "detail": "auto", + "prompt_cache_breakpoint": _BREAKPOINT, + }, + {"type": "input_file", "file_id": "file-abc", "prompt_cache_breakpoint": _BREAKPOINT}, + {"type": "input_text", "text": "unmarked extra text"}, + ], + }, + ] + +def _simple_expected_input(marker: str, *, marked: bool) -> list[JsonValue]: + text_block: Final = {"type": "input_text", "text": _prompt(marker, "user")} + return [ + { + "type": "message", + "role": "user", + "content": [{**text_block, **({"prompt_cache_breakpoint": _BREAKPOINT} if marked else {})}], + } + ] + +def _request_body(request: Request) -> dict[str, JsonValue]: + assert request.method == "POST" and request.target == "/v1/responses", request.target + return _JSON_OBJECT.validate_json(request.body) + +def _decoded_response_id(response_id: str) -> str: + decoded: Final = _RU._decode_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # reuse ID decoder + response_id + ) + raw_response_id: Final = decoded.get("response_id") + assert isinstance(raw_response_id, str), decoded + return raw_response_id + +def _spend_request_id_matches( + row: Mapping[str, JsonValue], + caller_response_id: str, + peer_response_id: str, + surface: _Surface, +) -> bool: + request_id: Final = row.get("request_id") + if not isinstance(request_id, str): + return False + match surface: + case "responses": + return request_id == caller_response_id + case "chat": + return _decoded_response_id(request_id) == peer_response_id + +def _spend_rows( + model: str, + caller_response_id: str, + peer_response_id: str, + surface: _Surface, +) -> tuple[dict[str, JsonValue], ...]: + def matching_rows(rows: list[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + row for row in rows if _spend_request_id_matches(row, caller_response_id, peer_response_id, surface) + ) + + rows: Final = eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda candidates: len(matching_rows(candidates)) == 1, + seconds=60, + ) + matched: Final = matching_rows(rows) + assert len(matched) == 1, matched + return matched + +def _response_id_from_chat_stream(text: str) -> str: + payloads: Final = tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: {") + ) + assert payloads, text + response_id: Final = payloads[0].get("id") + assert isinstance(response_id, str), payloads[0] + return response_id + +def _extra_body(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return { + key: value + for key, value in body.items() + if key not in {"model", "messages", "tools", "reasoning_effort", "stream", "num_retries"} + } + +def _sync_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: + base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" + model: Final = str(body["model"]) + messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) + tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) + extras: Final = _extra_body(body) + with OpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: + if stream: + response_stream: Final = client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=True, + extra_body=extras, + ) + chunks: Final = tuple(response_stream) + assert chunks + return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") + response: Final = client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=False, + extra_body=extras, + ) + return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") + +async def _async_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: + base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" + model: Final = str(body["model"]) + messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) + tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) + extras: Final = _extra_body(body) + async with AsyncOpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: + if stream: + response_stream: Final = await client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=True, + extra_body=extras, + ) + chunks: Final = tuple([chunk async for chunk in response_stream]) + assert chunks + return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") + response: Final = await client.chat.completions.create( + model=model, + messages=messages, + tools=tools, + reasoning_effort="low", + stream=False, + extra_body=extras, + ) + return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") + +async def _serve_chat( + gateway: Gateway, + body: dict[str, JsonValue], + client_kind: _ClientKind, + stream: bool, +) -> _Served: + match client_kind: + case "openai_sync": + return _sync_sdk_chat(gateway, body, stream) + case "openai_async": + return await _async_sdk_chat(gateway, body, stream) + case "httpx": + async with httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=20, + trust_env=False, + ) as client: + return await _raw_call( + client, + "/v1/chat/completions", + body, + _Call("chat", stream, _request_marker_from_body(body)), + ) + +def _request_marker_from_body(body: Mapping[str, JsonValue]) -> str: + match: Final = _MARKER.search(json.dumps(body).encode()) + assert match is not None, body + return match.group(1).decode() + +async def _raw_call( + client: httpx.AsyncClient, + path: str, + body: Mapping[str, JsonValue], + call: _Call, +) -> _Served: + async with client.stream( + "POST", + path, + json=body, + headers={"Authorization": f"Bearer {client.headers['Authorization'].removeprefix('Bearer ')}"}, + ) as response: + content: Final = await response.aread() + status: Final = response.status_code + text: Final = content.decode() + response_id: Final = ( + _response_id_from_chat_stream(text) + if status == 200 and call.surface == "chat" and call.stream + else _JSON_OBJECT.validate_json(content).get("id") + if status == 200 + else None + ) + return _Served(call, status, response_id if isinstance(response_id, str) else None, text) + +async def _send_call( + client: httpx.AsyncClient, + model: str, + call: _Call, +) -> _Served: + body: Final = ( + _simple_chat_body(model, call.marker, stream=call.stream) + if call.surface == "chat" + else { + "model": model, + "input": _simple_expected_input(call.marker, marked=True), + "stream": call.stream, + "num_retries": 0, + } + ) + path: Final = "/v1/chat/completions" if call.surface == "chat" else "/v1/responses" + try: + return await _raw_call(client, path, body, call) + except httpx.TransportError as error: + return _Served(call, 0, None, f"{type(error).__name__}: {error}") + +async def _burst( + base_url: str, + key: str, + model: str, + calls: tuple[_Call, ...], +) -> tuple[_Served, ...]: + async with httpx.AsyncClient( + base_url=base_url, + headers={"Authorization": f"Bearer {key}"}, + timeout=20, + trust_env=False, + limits=httpx.Limits(max_connections=100), + ) as client: + return tuple(await asyncio.gather(*(_send_call(client, model, call) for call in calls))) + +def _calls(count: int) -> tuple[_Call, ...]: + surfaces: Final[tuple[_Surface, ...]] = ("chat", "chat", "responses") + return tuple( + _Call( + surface=surfaces[index % len(surfaces)], + stream=index % 3 == 1, + marker=uuid.uuid4().hex, + ) + for index in range(count) + ) + +def _requests_for_marker(requests: tuple[Request, ...], marker: str) -> tuple[Request, ...]: + return tuple( + request + for request in requests + if request.method == "POST" and request.target == "/v1/responses" and _request_marker(request) == marker + ) + +def _peer_request_has_marker(request: Request, marker: str) -> bool: + body: Final = _request_body(request) + return _contains_breakpoint(body) and _request_marker(request) == marker + +def _peer_marker_matches_response(served: _Served, requests: tuple[Request, ...]) -> bool: + peer_requests: Final = _requests_for_marker(requests, served.call.marker) + assert len(peer_requests) == 1, (served, peer_requests) + (peer_request,) = peer_requests + return _peer_request_has_marker(peer_request, served.call.marker) + +def _assert_spend_for_result(served: _Served, model: str) -> None: + assert served.status == 200 and served.response_id is not None, served + peer_response_id: Final = _response_id(served.call.marker) + match served.call.surface: + case "responses": + (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) + case "chat": + assert _decoded_response_id(served.response_id) == peer_response_id, served + (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) + request_id: Final = row.get("request_id") + assert isinstance(request_id, str), row + match served.call.surface: + case "responses": + assert request_id == served.response_id, row + case "chat": + assert _decoded_response_id(request_id) == served.response_id, row diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py index 75d8ffac9f7..9332b4f3507 100644 --- a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint.py @@ -1,544 +1,33 @@ from __future__ import annotations -import asyncio -import json -import re -import signal -import threading import uuid -from collections.abc import Iterable, Mapping -from contextlib import ExitStack -from dataclasses import dataclass -from pathlib import Path -from queue import SimpleQueue -from typing import Final, Literal, TypeAlias, cast -from urllib.parse import urlsplit +from typing import Final import httpx -import psutil import pytest -import yaml -from openai import AsyncOpenAI, OpenAI -from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolUnionParam -from pydantic import JsonValue, TypeAdapter - -from integration._support.client import Gateway, eventually -from integration._support.database import read_rows -from integration._support.process import owned_proxy_process -from integration._support.wire import Reply, Request, Wire, wire_server -from litellm.responses.utils import ResponsesAPIRequestUtils as _RU - -_MODEL: Final = "openai/gpt-5.6" -_UNSUPPORTED_MODEL: Final = "openai/gpt-5.4-mini" -_CONFIG_MODEL: Final = "responses-bridge-cache-breakpoint-chaos" -_API_KEY: Final = "synthetic-responses-bridge-key" -_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})") -_STARTED_WORKER: Final[re.Pattern[str]] = re.compile(r"Started server process \[(\d+)\]") -_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) -_BREAKPOINT: Final[dict[str, JsonValue]] = {"mode": "explicit"} -_TOOLS: Final[list[JsonValue]] = [ - { - "type": "function", - "function": { - "name": "synthetic_tool", - "description": "Synthetic bridge test tool", - "parameters": {"type": "object", "properties": {}}, - }, - } -] -_IMAGE_URL: Final = "data:image/png;base64,aGVsbG8=" -_ClientKind: TypeAlias = Literal["openai_sync", "openai_async", "httpx"] -_Surface: TypeAlias = Literal["chat", "responses"] - - -@dataclass(frozen=True, slots=True) -class _Call: - surface: _Surface - stream: bool - marker: str - - -@dataclass(frozen=True, slots=True) -class _Served: - call: _Call - status: int - response_id: str | None - text: str - - -def _response_id(marker: str) -> str: - return f"resp_{marker}" - - -def _request_marker(request: Request) -> str: - match: Final = _MARKER.search(request.body) - assert match is not None, request.body - return match.group(1).decode() - - -def _contains_breakpoint(value: JsonValue) -> bool: - if isinstance(value, dict): - return "prompt_cache_breakpoint" in value or any(_contains_breakpoint(item) for item in value.values()) - if isinstance(value, list): - return any(_contains_breakpoint(item) for item in value) - return False - - -def _responses_body(marker: str) -> dict[str, JsonValue]: - response_id: Final = _response_id(marker) - return _JSON_OBJECT.validate_python( - { - "id": response_id, - "object": "response", - "created_at": 1, - "status": "completed", - "model": "gpt-5.6", - "output": [ - { - "id": f"msg_{marker}", - "type": "message", - "role": "assistant", - "status": "completed", - "content": [{"type": "output_text", "text": f"answer marker-{marker}", "annotations": []}], - } - ], - "usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}, - } - ) - - -def _responses_reply(request: Request, *, reject_breakpoints: bool = False) -> Reply: - if request.method == "GET" and request.target == "/v1/models": - return Reply(body=b'{"object":"list","data":[{"id":"gpt-5.6","object":"model"}]}') - body: Final = _JSON_OBJECT.validate_json(request.body) - if reject_breakpoints and _contains_breakpoint(body): - return Reply( - status=400, - body=json.dumps( - { - "error": { - "message": "prompt_cache_breakpoint is not supported on this model", - "type": "invalid_request_error", - "param": None, - "code": None, - } - } - ).encode(), - ) - marker: Final = _request_marker(request) - stream: Final = body.get("stream") is True - response: Final = _responses_body(marker) - if not stream: - return Reply(body=json.dumps(response).encode()) - created: Final = { - "type": "response.created", - "sequence_number": 0, - "response": {**response, "status": "in_progress", "output": []}, - } - delta: Final = { - "type": "response.output_text.delta", - "sequence_number": 1, - "item_id": f"msg_{marker}", - "output_index": 0, - "content_index": 0, - "delta": f"answer marker-{marker}", - } - completed: Final = {"type": "response.completed", "sequence_number": 2, "response": response} - events: Final = (created, delta, completed) - return Reply( - content_type="text/event-stream", - chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), - ) - - -def _prompt(marker: str, label: str) -> str: - return f"{label} marker-{marker}" - - -def _simple_chat_body( - model: str, - marker: str, - *, - stream: bool = False, - marked: bool = True, - system_as_string: bool = False, - prompt_cache_options: dict[str, JsonValue] | None = None, -) -> dict[str, JsonValue]: - marker_field: Final = {"prompt_cache_breakpoint": _BREAKPOINT} if marked else {} - user: Final = [{"type": "text", "text": _prompt(marker, "user"), **marker_field}] - messages: Final = ( - [{"role": "system", "content": _prompt(marker, "system")}, {"role": "user", "content": user}] - if system_as_string - else [{"role": "user", "content": user}] - ) - return _JSON_OBJECT.validate_python( - { - "model": model, - "messages": messages, - "tools": _TOOLS, - "reasoning_effort": "low", - "stream": stream, - "num_retries": 0, - **({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}), - } - ) - - -def _multimodal_chat_body(model: str, marker: str, stream: bool) -> dict[str, JsonValue]: - return _JSON_OBJECT.validate_python( - { - "model": model, - "messages": [ - { - "role": "system", - "content": [ - { - "type": "text", - "text": _prompt(marker, "system"), - "prompt_cache_breakpoint": _BREAKPOINT, - } - ], - }, - { - "role": "user", - "content": [ - {"type": "text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, - { - "type": "image_url", - "image_url": {"url": _IMAGE_URL}, - "prompt_cache_breakpoint": _BREAKPOINT, - }, - { - "type": "file", - "file": {"file_id": "file-abc"}, - "prompt_cache_breakpoint": _BREAKPOINT, - }, - {"type": "text", "text": "unmarked extra text"}, - ], - }, - ], - "tools": _TOOLS, - "reasoning_effort": "low", - "stream": stream, - "num_retries": 0, - "prompt_cache_options": {"mode": "explicit"}, - } - ) - - -def _expected_multimodal_input(marker: str) -> list[JsonValue]: - return [ - { - "type": "message", - "role": "system", - "content": [ - {"type": "input_text", "text": _prompt(marker, "system"), "prompt_cache_breakpoint": _BREAKPOINT} - ], - }, - { - "type": "message", - "role": "user", - "content": [ - {"type": "input_text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT}, - { - "type": "input_image", - "image_url": _IMAGE_URL, - "detail": "auto", - "prompt_cache_breakpoint": _BREAKPOINT, - }, - {"type": "input_file", "file_id": "file-abc", "prompt_cache_breakpoint": _BREAKPOINT}, - {"type": "input_text", "text": "unmarked extra text"}, - ], - }, - ] - - -def _simple_expected_input(marker: str, *, marked: bool) -> list[JsonValue]: - text_block: Final = {"type": "input_text", "text": _prompt(marker, "user")} - return [ - { - "type": "message", - "role": "user", - "content": [{**text_block, **({"prompt_cache_breakpoint": _BREAKPOINT} if marked else {})}], - } - ] - - -def _request_body(request: Request) -> dict[str, JsonValue]: - assert request.method == "POST" and request.target == "/v1/responses", request.target - return _JSON_OBJECT.validate_json(request.body) - - -def _decoded_response_id(response_id: str) -> str: - decoded: Final = _RU._decode_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # reuse ID decoder - response_id - ) - raw_response_id: Final = decoded.get("response_id") - assert isinstance(raw_response_id, str), decoded - return raw_response_id - - -def _spend_request_id_matches( - row: Mapping[str, JsonValue], - caller_response_id: str, - peer_response_id: str, - surface: _Surface, -) -> bool: - request_id: Final = row.get("request_id") - if not isinstance(request_id, str): - return False - match surface: - case "responses": - return request_id == caller_response_id - case "chat": - return _decoded_response_id(request_id) == peer_response_id - - -def _spend_rows( - model: str, - caller_response_id: str, - peer_response_id: str, - surface: _Surface, -) -> tuple[dict[str, JsonValue], ...]: - def matching_rows(rows: list[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]: - return tuple( - row for row in rows if _spend_request_id_matches(row, caller_response_id, peer_response_id, surface) - ) - - rows: Final = eventually( - lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), - lambda candidates: len(matching_rows(candidates)) == 1, - seconds=60, - ) - matched: Final = matching_rows(rows) - assert len(matched) == 1, matched - return matched - - -def _response_id_from_chat_stream(text: str) -> str: - payloads: Final = tuple( - _JSON_OBJECT.validate_json(line.removeprefix("data: ")) - for line in text.splitlines() - if line.startswith("data: {") - ) - assert payloads, text - response_id: Final = payloads[0].get("id") - assert isinstance(response_id, str), payloads[0] - return response_id - - -def _extra_body(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: - return { - key: value - for key, value in body.items() - if key not in {"model", "messages", "tools", "reasoning_effort", "stream", "num_retries"} - } - - -def _sync_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: - base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" - model: Final = str(body["model"]) - messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) - tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) - extras: Final = _extra_body(body) - with OpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: - if stream: - response_stream: Final = client.chat.completions.create( - model=model, - messages=messages, - tools=tools, - reasoning_effort="low", - stream=True, - extra_body=extras, - ) - chunks: Final = tuple(response_stream) - assert chunks - return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") - response: Final = client.chat.completions.create( - model=model, - messages=messages, - tools=tools, - reasoning_effort="low", - stream=False, - extra_body=extras, - ) - return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") - - -async def _async_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served: - base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1" - model: Final = str(body["model"]) - messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"]) - tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"]) - extras: Final = _extra_body(body) - async with AsyncOpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client: - if stream: - response_stream: Final = await client.chat.completions.create( - model=model, - messages=messages, - tools=tools, - reasoning_effort="low", - stream=True, - extra_body=extras, - ) - chunks: Final = tuple([chunk async for chunk in response_stream]) - assert chunks - return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "") - response: Final = await client.chat.completions.create( - model=model, - messages=messages, - tools=tools, - reasoning_effort="low", - stream=False, - extra_body=extras, - ) - return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "") - - -async def _serve_chat( - gateway: Gateway, - body: dict[str, JsonValue], - client_kind: _ClientKind, - stream: bool, -) -> _Served: - match client_kind: - case "openai_sync": - return _sync_sdk_chat(gateway, body, stream) - case "openai_async": - return await _async_sdk_chat(gateway, body, stream) - case "httpx": - async with httpx.AsyncClient( - base_url=str(gateway.client.base_url), - headers={"Authorization": f"Bearer {gateway.key}"}, - timeout=20, - trust_env=False, - ) as client: - return await _raw_call( - client, - "/v1/chat/completions", - body, - _Call("chat", stream, _request_marker_from_body(body)), - ) - - -def _request_marker_from_body(body: Mapping[str, JsonValue]) -> str: - match: Final = _MARKER.search(json.dumps(body).encode()) - assert match is not None, body - return match.group(1).decode() - - -async def _raw_call( - client: httpx.AsyncClient, - path: str, - body: Mapping[str, JsonValue], - call: _Call, -) -> _Served: - async with client.stream( - "POST", - path, - json=body, - headers={"Authorization": f"Bearer {client.headers['Authorization'].removeprefix('Bearer ')}"}, - ) as response: - content: Final = await response.aread() - status: Final = response.status_code - text: Final = content.decode() - response_id: Final = ( - _response_id_from_chat_stream(text) - if status == 200 and call.surface == "chat" and call.stream - else _JSON_OBJECT.validate_json(content).get("id") - if status == 200 - else None - ) - return _Served(call, status, response_id if isinstance(response_id, str) else None, text) - - -async def _send_call( - client: httpx.AsyncClient, - model: str, - call: _Call, -) -> _Served: - body: Final = ( - _simple_chat_body(model, call.marker, stream=call.stream) - if call.surface == "chat" - else { - "model": model, - "input": _simple_expected_input(call.marker, marked=True), - "stream": call.stream, - "num_retries": 0, - } - ) - path: Final = "/v1/chat/completions" if call.surface == "chat" else "/v1/responses" - try: - return await _raw_call(client, path, body, call) - except httpx.TransportError as error: - return _Served(call, 0, None, f"{type(error).__name__}: {error}") - - -async def _burst( - base_url: str, - key: str, - model: str, - calls: tuple[_Call, ...], -) -> tuple[_Served, ...]: - async with httpx.AsyncClient( - base_url=base_url, - headers={"Authorization": f"Bearer {key}"}, - timeout=20, - trust_env=False, - limits=httpx.Limits(max_connections=100), - ) as client: - return tuple(await asyncio.gather(*(_send_call(client, model, call) for call in calls))) - - -def _calls(count: int) -> tuple[_Call, ...]: - surfaces: Final[tuple[_Surface, ...]] = ("chat", "chat", "responses") - return tuple( - _Call( - surface=surfaces[index % len(surfaces)], - stream=index % 3 == 1, - marker=uuid.uuid4().hex, - ) - for index in range(count) - ) - - -def _requests_for_marker(requests: tuple[Request, ...], marker: str) -> tuple[Request, ...]: - return tuple( - request - for request in requests - if request.method == "POST" and request.target == "/v1/responses" and _request_marker(request) == marker - ) - - -def _peer_request_has_marker(request: Request, marker: str) -> bool: - body: Final = _request_body(request) - return _contains_breakpoint(body) and _request_marker(request) == marker - - -def _peer_marker_matches_response(served: _Served, requests: tuple[Request, ...]) -> bool: - peer_requests: Final = _requests_for_marker(requests, served.call.marker) - assert len(peer_requests) == 1, (served, peer_requests) - (peer_request,) = peer_requests - return _peer_request_has_marker(peer_request, served.call.marker) - - -def _assert_spend_for_result(served: _Served, model: str) -> None: - assert served.status == 200 and served.response_id is not None, served - peer_response_id: Final = _response_id(served.call.marker) - match served.call.surface: - case "responses": - (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) - case "chat": - assert _decoded_response_id(served.response_id) == peer_response_id, served - (row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface) - request_id: Final = row.get("request_id") - assert isinstance(request_id, str), row - match served.call.surface: - case "responses": - assert request_id == served.response_id, row - case "chat": - assert _decoded_response_id(request_id) == served.response_id, row +from pydantic import JsonValue +from integration._support.client import Gateway +from integration._support.wire import wire_server +from integration.providers._responses_bridge_prompt_cache_breakpoint import ( + _BREAKPOINT, + _Call, + _ClientKind, + _JSON_OBJECT, + _MODEL, + _UNSUPPORTED_MODEL, + _assert_spend_for_result, + _contains_breakpoint, + _expected_multimodal_input, + _multimodal_chat_body, + _prompt, + _raw_call, + _request_body, + _responses_reply, + _serve_chat, + _simple_chat_body, + _simple_expected_input, +) @pytest.mark.parametrize("client_kind", ("openai_sync", "openai_async", "httpx")) @pytest.mark.parametrize("stream", (False, True)) @@ -559,7 +48,6 @@ async def test_caller_prompt_cache_breakpoints_survive_chat_to_responses_bridge( assert peer_body["input"] == _expected_multimodal_input(marker), peer_body assert peer_body["prompt_cache_options"] == {"mode": "explicit"}, peer_body - def _expected_uninjected_system_bridge_body( marker: str, prompt_cache_options: dict[str, JsonValue] | None = None, @@ -582,7 +70,6 @@ def _expected_uninjected_system_bridge_body( **({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}), } - async def test_deployment_cache_control_injection_without_options_is_unchanged(gateway: Gateway) -> None: marker: Final = uuid.uuid4().hex with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: @@ -607,7 +94,6 @@ async def test_deployment_cache_control_injection_without_options_is_unchanged(g assert not _contains_breakpoint(peer_body), peer_body assert "prompt_cache_options" not in peer_body, peer_body - async def test_deployment_prompt_cache_options_override_is_unchanged(gateway: Gateway) -> None: marker: Final = uuid.uuid4().hex options: Final[dict[str, JsonValue]] = {"mode": "implicit", "ttl": "30m"} @@ -632,7 +118,6 @@ async def test_deployment_prompt_cache_options_override_is_unchanged(gateway: Ga assert peer_body == _expected_uninjected_system_bridge_body(marker, options), peer_body assert not _contains_breakpoint(peer_body), peer_body - async def test_unmarked_bridge_and_direct_responses_marker_are_forwarded_unchanged(gateway: Gateway) -> None: marker: Final = uuid.uuid4().hex with wire_server(_responses_reply) as wire, gateway.scenario() as scenario: @@ -689,7 +174,6 @@ async def test_unmarked_bridge_and_direct_responses_marker_are_forwarded_unchang (direct_peer,) = wire.drain() assert _request_body(direct_peer)["input"] == direct_input, _request_body(direct_peer) - @pytest.mark.parametrize("stream", (False, True)) async def test_unsupported_model_drops_breakpoints_without_rejecting_the_request( gateway: Gateway, @@ -715,197 +199,3 @@ async def test_unsupported_model_drops_breakpoints_without_rejecting_the_request peer_body: Final = _request_body(peer_request) assert peer_body["input"] == _simple_expected_input(marker, marked=False), peer_body assert not _contains_breakpoint(peer_body), peer_body - - -def _chaos_config(wire: Wire, tmp_path: Path) -> Path: - base_config: Final = _JSON_OBJECT.validate_python( - yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) - ) - config: Final = { - **base_config, - "model_list": [ - { - "model_name": _CONFIG_MODEL, - "litellm_params": { - "model": _MODEL, - "api_base": wire.url + "/v1", - "api_key": _API_KEY, - }, - }, - ], - } - path: Final = tmp_path / "responses-bridge-cache-breakpoint-chaos.yaml" - path.write_text(yaml.safe_dump(config)) - return path - - -def _open_upstream_connections(pid: int, port: int) -> int: - return sum( - 1 - for connection in psutil.Process(pid).net_connections(kind="tcp") - if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port - ) - - -@pytest.mark.timeout(180) -async def test_worker_and_peer_outages_preserve_markers_and_recover( - gateway: Gateway, - tmp_path: Path, -) -> None: - calls: Final = _calls(30) - release: Final = threading.Event() - early_release: Final = threading.Event() - outage_release: Final = threading.Event() - held_markers: Final[SimpleQueue[str]] = SimpleQueue() - early_calls: Final = calls[:10] - early_markers: Final = frozenset(call.marker for call in early_calls) - - def held(request: Request) -> Reply: - if request.method == "GET" and request.target == "/v1/models": - return _responses_reply(request) - marker: Final = _request_marker(request) - held_markers.put(marker) - gate: Final = early_release if marker in early_markers else release - assert gate.wait(timeout=60), "The worker-kill burst was never released" - return _responses_reply(request) - - with ExitStack() as peer_stack: - wire: Final = peer_stack.enter_context(wire_server(held)) - config: Final = _chaos_config(wire, tmp_path) - with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: - try: - candidate: Final = owned.gateway - workers: Final[tuple[int, ...]] = eventually( - lambda: tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(owned.log.read_text())), - lambda pids: len(pids) == 2, - seconds=30, - ) - async with httpx.AsyncClient( - base_url=str(candidate.client.base_url), - headers={"Authorization": f"Bearer {candidate.key}"}, - timeout=20, - trust_env=False, - limits=httpx.Limits(max_connections=100), - ) as client: - burst_tasks: Final = tuple( - asyncio.create_task(_send_call(client, _CONFIG_MODEL, call)) for call in calls - ) - await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == len(calls), 60) - early_release.set() - early_served: Final = await asyncio.gather(*burst_tasks[: len(early_calls)]) - early_successful: Final = tuple(item for item in early_served if item.status == 200) - for item in early_successful: - _assert_spend_for_result(item, _CONFIG_MODEL) - upstream_port_value: Final = urlsplit(wire.url).port - assert upstream_port_value is not None - upstream_port: Final = upstream_port_value - active_by_worker: Final = eventually( - lambda: {pid: _open_upstream_connections(pid, upstream_port) for pid in workers}, - lambda counts: sum(counts.values()) == len(calls) - len(early_calls), - seconds=30, - ) - victim_pid: Final = max(workers, key=active_by_worker.__getitem__) - survivor_pids: Final = tuple(pid for pid in workers if pid != victim_pid) - assert active_by_worker[victim_pid] > 0 and len(survivor_pids) == 1, active_by_worker - (survivor_pid,) = survivor_pids - victim: Final = psutil.Process(victim_pid) - victim.suspend() - victim.send_signal(signal.SIGKILL) - release.set() - remaining_served: Final = await asyncio.gather(*burst_tasks[len(early_calls) :]) - served: Final = (*early_served, *remaining_served) - successful: Final = tuple(item for item in served if item.status == 200) - connection_errors: Final = tuple(item for item in served if item.status == 0) - print(f"worker-kill burst: {len(successful)} HTTP 200, {len(connection_errors)} connection errors") - assert len(successful) + len(connection_errors) == len(calls), { - "successes": len(successful), - "connection_errors": len(connection_errors), - "responses": served, - } - assert successful and connection_errors, { - "successes": len(successful), - "connection_errors": len(connection_errors), - } - follow_ups: Final = ( - _Call("chat", False, uuid.uuid4().hex), - _Call("responses", False, uuid.uuid4().hex), - ) - recovered: Final = await _burst( - str(candidate.client.base_url), - candidate.key, - _CONFIG_MODEL, - follow_ups, - ) - assert all(item.status == 200 for item in recovered), recovered - assert psutil.pid_exists(survivor_pid), survivor_pid - received_after_worker_kill: Final = wire.drain() - worker_marker_failures: Final = tuple( - item.call.marker - for item in (*successful, *recovered) - if not _peer_marker_matches_response(item, received_after_worker_kill) - ) - for item in (*successful, *recovered): - assert item.response_id is not None, item - _assert_spend_for_result(item, _CONFIG_MODEL) - - peer_stack.close() - outage_seen: Final[SimpleQueue[str]] = SimpleQueue() - - def outage(request: Request) -> Reply: - if request.method == "GET" and request.target == "/v1/models": - return _responses_reply(request) - outage_seen.put(_request_marker(request)) - assert outage_release.wait(timeout=60), "The peer-outage burst was never stopped" - return Reply( - status=503, - body=b'{"error":{"message":"synthetic peer outage","type":"server_error"}}', - ) - - peer_stack.enter_context(wire_server(outage, port=upstream_port)) - outage_calls: Final = _calls(12) - outage_burst: Final = asyncio.create_task( - _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, outage_calls) - ) - await asyncio.to_thread(eventually, outage_seen.qsize, lambda size: size == len(outage_calls), 30) - outage_release.set() - peer_stack.close() - outage_served: Final = await outage_burst - assert len(outage_served) == len(outage_calls), outage_served - assert all(item.status >= 400 and "error" in item.text.lower() for item in outage_served), outage_served - down_call: Final = _Call("chat", False, uuid.uuid4().hex) - (down_response,) = await _burst( - str(candidate.client.base_url), - candidate.key, - _CONFIG_MODEL, - (down_call,), - ) - assert down_response.status >= 400 and "error" in down_response.text.lower(), down_response - - restarted_wire: Final = peer_stack.enter_context(wire_server(_responses_reply, port=upstream_port)) - recovery_calls: Final = ( - _Call("chat", False, uuid.uuid4().hex), - _Call("responses", False, uuid.uuid4().hex), - ) - recovered_after_peer_restart: Final = await _burst( - str(candidate.client.base_url), - candidate.key, - _CONFIG_MODEL, - recovery_calls, - ) - assert all(item.status == 200 for item in recovered_after_peer_restart), recovered_after_peer_restart - restarted_requests: Final = restarted_wire.drain() - recovery_marker_failures: Final = tuple( - item.call.marker - for item in recovered_after_peer_restart - if not _peer_marker_matches_response(item, restarted_requests) - ) - for item in recovered_after_peer_restart: - assert item.response_id is not None, item - _assert_spend_for_result(item, _CONFIG_MODEL) - assert not (*worker_marker_failures, *recovery_marker_failures), { - "worker_marker_failures": worker_marker_failures, - "recovery_marker_failures": recovery_marker_failures, - } - finally: - release.set() - outage_release.set() diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_chaos.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_chaos.py new file mode 100644 index 00000000000..10d5cfd9646 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_chaos.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +import asyncio +import re +import signal +import threading +import uuid +from contextlib import ExitStack +from pathlib import Path +from queue import SimpleQueue +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.providers._responses_bridge_prompt_cache_breakpoint import ( + _Call, + _JSON_OBJECT, + _MODEL, + _assert_spend_for_result, + _burst, + _calls, + _peer_marker_matches_response, + _request_marker, + _responses_reply, + _send_call, +) + +_CONFIG_MODEL: Final = "responses-bridge-cache-breakpoint-chaos" + +_API_KEY: Final = "synthetic-responses-bridge-key" + +_STARTED_WORKER: Final[re.Pattern[str]] = re.compile(r"Started server process \[(\d+)\]") + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + base_config: Final = _JSON_OBJECT.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + config: Final = { + **base_config, + "model_list": [ + { + "model_name": _CONFIG_MODEL, + "litellm_params": { + "model": _MODEL, + "api_base": wire.url + "/v1", + "api_key": _API_KEY, + }, + }, + ], + } + path: Final = tmp_path / "responses-bridge-cache-breakpoint-chaos.yaml" + path.write_text(yaml.safe_dump(config)) + return path + +def _open_upstream_connections(pid: int, port: int) -> int: + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + +@pytest.mark.timeout(180) +async def test_worker_and_peer_outages_preserve_markers_and_recover( + gateway: Gateway, + tmp_path: Path, +) -> None: + calls: Final = _calls(30) + release: Final = threading.Event() + early_release: Final = threading.Event() + outage_release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + early_calls: Final = calls[:10] + early_markers: Final = frozenset(call.marker for call in early_calls) + + def held(request: Request) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return _responses_reply(request) + marker: Final = _request_marker(request) + held_markers.put(marker) + gate: Final = early_release if marker in early_markers else release + assert gate.wait(timeout=60), "The worker-kill burst was never released" + return _responses_reply(request) + + with ExitStack() as peer_stack: + wire: Final = peer_stack.enter_context(wire_server(held)) + config: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + try: + candidate: Final = owned.gateway + workers: Final[tuple[int, ...]] = eventually( + lambda: tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + async with httpx.AsyncClient( + base_url=str(candidate.client.base_url), + headers={"Authorization": f"Bearer {candidate.key}"}, + timeout=20, + trust_env=False, + limits=httpx.Limits(max_connections=100), + ) as client: + burst_tasks: Final = tuple( + asyncio.create_task(_send_call(client, _CONFIG_MODEL, call)) for call in calls + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == len(calls), 60) + early_release.set() + early_served: Final = await asyncio.gather(*burst_tasks[: len(early_calls)]) + early_successful: Final = tuple(item for item in early_served if item.status == 200) + for item in early_successful: + _assert_spend_for_result(item, _CONFIG_MODEL) + upstream_port_value: Final = urlsplit(wire.url).port + assert upstream_port_value is not None + upstream_port: Final = upstream_port_value + active_by_worker: Final = eventually( + lambda: {pid: _open_upstream_connections(pid, upstream_port) for pid in workers}, + lambda counts: sum(counts.values()) == len(calls) - len(early_calls), + seconds=30, + ) + victim_pid: Final = max(workers, key=active_by_worker.__getitem__) + survivor_pids: Final = tuple(pid for pid in workers if pid != victim_pid) + assert active_by_worker[victim_pid] > 0 and len(survivor_pids) == 1, active_by_worker + (survivor_pid,) = survivor_pids + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + remaining_served: Final = await asyncio.gather(*burst_tasks[len(early_calls) :]) + served: Final = (*early_served, *remaining_served) + successful: Final = tuple(item for item in served if item.status == 200) + connection_errors: Final = tuple(item for item in served if item.status == 0) + print(f"worker-kill burst: {len(successful)} HTTP 200, {len(connection_errors)} connection errors") + assert len(successful) + len(connection_errors) == len(calls), { + "successes": len(successful), + "connection_errors": len(connection_errors), + "responses": served, + } + assert successful and connection_errors, { + "successes": len(successful), + "connection_errors": len(connection_errors), + } + follow_ups: Final = ( + _Call("chat", False, uuid.uuid4().hex), + _Call("responses", False, uuid.uuid4().hex), + ) + recovered: Final = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + follow_ups, + ) + assert all(item.status == 200 for item in recovered), recovered + assert psutil.pid_exists(survivor_pid), survivor_pid + received_after_worker_kill: Final = wire.drain() + worker_marker_failures: Final = tuple( + item.call.marker + for item in (*successful, *recovered) + if not _peer_marker_matches_response(item, received_after_worker_kill) + ) + for item in (*successful, *recovered): + assert item.response_id is not None, item + _assert_spend_for_result(item, _CONFIG_MODEL) + + peer_stack.close() + outage_seen: Final[SimpleQueue[str]] = SimpleQueue() + + def outage(request: Request) -> Reply: + if request.method == "GET" and request.target == "/v1/models": + return _responses_reply(request) + outage_seen.put(_request_marker(request)) + assert outage_release.wait(timeout=60), "The peer-outage burst was never stopped" + return Reply( + status=503, + body=b'{"error":{"message":"synthetic peer outage","type":"server_error"}}', + ) + + peer_stack.enter_context(wire_server(outage, port=upstream_port)) + outage_calls: Final = _calls(12) + outage_burst: Final = asyncio.create_task( + _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, outage_calls) + ) + await asyncio.to_thread(eventually, outage_seen.qsize, lambda size: size == len(outage_calls), 30) + outage_release.set() + peer_stack.close() + outage_served: Final = await outage_burst + assert len(outage_served) == len(outage_calls), outage_served + assert all(item.status >= 400 and "error" in item.text.lower() for item in outage_served), outage_served + down_call: Final = _Call("chat", False, uuid.uuid4().hex) + (down_response,) = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + (down_call,), + ) + assert down_response.status >= 400 and "error" in down_response.text.lower(), down_response + + restarted_wire: Final = peer_stack.enter_context(wire_server(_responses_reply, port=upstream_port)) + recovery_calls: Final = ( + _Call("chat", False, uuid.uuid4().hex), + _Call("responses", False, uuid.uuid4().hex), + ) + recovered_after_peer_restart: Final = await _burst( + str(candidate.client.base_url), + candidate.key, + _CONFIG_MODEL, + recovery_calls, + ) + assert all(item.status == 200 for item in recovered_after_peer_restart), recovered_after_peer_restart + restarted_requests: Final = restarted_wire.drain() + recovery_marker_failures: Final = tuple( + item.call.marker + for item in recovered_after_peer_restart + if not _peer_marker_matches_response(item, restarted_requests) + ) + for item in recovered_after_peer_restart: + assert item.response_id is not None, item + _assert_spend_for_result(item, _CONFIG_MODEL) + assert not (*worker_marker_failures, *recovery_marker_failures), { + "worker_marker_failures": worker_marker_failures, + "recovery_marker_failures": recovery_marker_failures, + } + finally: + release.set() + outage_release.set() 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 c6a4cdaed39..612232e1692 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 @@ -4291,6 +4291,42 @@ def test_prompt_cache_breakpoints_are_dropped_for_unsupported_models() -> None: ] +@pytest.mark.parametrize( + ("litellm_params", "keep_marker"), + (({"base_model": "gpt-5.6"}, True), ({}, False)), + ids=("supported-base-model", "missing-base-model"), +) +def test_prompt_cache_breakpoint_supports_model_alias_with_base_model( + litellm_params: dict[str, object], + keep_marker: bool, +) -> None: + handler: Final = LiteLLMResponsesTransformationHandler() + cache_breakpoint: Final = {"mode": "explicit"} + marked_content: Final = {"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": cache_breakpoint} + + request: Final = handler.transform_request( + model="mydeployment", + messages=[{"role": "user", "content": [marked_content]}], + optional_params={}, + litellm_params=litellm_params, + headers={}, + litellm_logging_obj=Mock(), + ) + + expected_content: Final = { + "type": "input_text", + "text": "Stable prefix", + **({"prompt_cache_breakpoint": cache_breakpoint} if keep_marker else {}), + } + assert request["input"] == [ + { + "type": "message", + "role": "user", + "content": [expected_content], + } + ] + + def test_mid_conversation_system_string_stays_in_input_after_a_user_turn(): handler: Final = LiteLLMResponsesTransformationHandler() From c8720df299af84faddd99a9c6bad20a4ca927e96 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 01:54:18 +0000 Subject: [PATCH 3/9] 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)), From 202c4fbed29609f2e3fab019ba17844c2f1e5094 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 02:21:19 +0000 Subject: [PATCH 4/9] fix(responses): justify bridge stripping casts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_responses_transformation/transformation.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 96db88475ea..367c5b31321 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -84,16 +84,16 @@ _CHAT_CONTENT_ITEM: Final = TypeAdapter(dict[str, object]) 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) + content_block: Final = cast(dict[str, object], value) # cast-ok: isinstance confirms the content block is a mapping 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): - list_content: Final = cast(list[object], value) + list_content: Final = cast(list[object], value) # cast-ok: isinstance confirms a list of content blocks return [_strip_prompt_cache_breakpoint_from_content_block(item) for item in list_content] if isinstance(value, tuple): - tuple_content: Final = cast(tuple[object, ...], value) + tuple_content: Final = cast(tuple[object, ...], value) # cast-ok: isinstance confirms a tuple of content blocks return tuple(_strip_prompt_cache_breakpoint_from_content_block(item) for item in tuple_content) return _strip_prompt_cache_breakpoint_from_content_block(value) @@ -101,7 +101,7 @@ def _strip_prompt_cache_breakpoints_from_content(value: object) -> object: 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) + input_item: Final = cast(dict[str, object], value) # cast-ok: isinstance confirms a Responses input item mapping return { key: _strip_prompt_cache_breakpoints_from_content(item) if key in ("content", "output") else item for key, item in input_item.items() From ccf03ee7dfa096a04d6106375478010e16a25543 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 03:13:32 +0000 Subject: [PATCH 5/9] fix(responses): keep cache breakpoints out of non-bridge converter callers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../transformation.py | 28 +++-- ...responses_transformation_transformation.py | 115 ++++++++++++++++++ 2 files changed, 134 insertions(+), 9 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 367c5b31321..d69b2915067 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -393,6 +393,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return None, index def convert_chat_completion_messages_to_responses_api( + self, + messages: list["AllMessageValues"], + *, + keep_prompt_cache_breakpoints: bool = False, + ) -> tuple[list[object], str | None]: + converted_input_items, instructions = self._convert_chat_completion_messages_to_responses_input(messages) + return ( + converted_input_items + if keep_prompt_cache_breakpoints + else _strip_prompt_cache_breakpoints(converted_input_items), + instructions, + ) + + def _convert_chat_completion_messages_to_responses_input( self, messages: list["AllMessageValues"] ) -> tuple[list[object], str | None]: input_items: Final[list[object]] = [] @@ -623,23 +637,19 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): litellm_logging_obj: "LiteLLMLoggingObj", client: object | None = None, ) -> dict: - converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(messages) base_model: Final = litellm_params.get("base_model") supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) or ( isinstance(base_model, str) and bool(base_model) and supports_openai_prompt_cache_breakpoint(base_model) ) - input_items_without_unsupported_markers: Final = ( - converted_input_items - if supports_prompt_cache_breakpoint - else _strip_prompt_cache_breakpoints(converted_input_items) + converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api( + messages, + keep_prompt_cache_breakpoints=supports_prompt_cache_breakpoint, ) # OpenAI's Responses API rejects an empty input. For a system-only # request, carry the system message as a system-role input item instead # of instructions, mirroring how non-string system content is already # handled in convert_chat_completion_messages_to_responses_api. - is_system_only_request: Final = ( - not input_items_without_unsupported_markers and converted_instructions is not None - ) + is_system_only_request: Final = not converted_input_items and converted_instructions is not None input_items: Final = ( [ { @@ -649,7 +659,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): } ] if is_system_only_request - else input_items_without_unsupported_markers + else converted_input_items ) instructions: Final = None if is_system_only_request else converted_instructions 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 a45a56a5226..fcb44958ffb 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 @@ -1,3 +1,4 @@ +import copy import datetime import json import os @@ -4333,6 +4334,120 @@ def test_prompt_cache_breakpoints_are_dropped_from_function_call_output_for_unsu ] +def test_convert_chat_completion_messages_to_responses_api_drops_prompt_cache_breakpoints_unless_kept() -> None: + handler: Final = LiteLLMResponsesTransformationHandler() + cache_breakpoint: Final = {"mode": "explicit"} + image_data_url: Final = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg==" + file_data: Final = "data:application/pdf;base64,JVBERi0xLjQK" + messages: Final = cast( + list[AllMessageValues], + [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Review these inputs", "prompt_cache_breakpoint": cache_breakpoint}, + { + "type": "image_url", + "image_url": {"url": image_data_url}, + "prompt_cache_breakpoint": cache_breakpoint, + }, + { + "type": "file", + "file": {"file_data": file_data, "filename": "input.pdf"}, + "prompt_cache_breakpoint": cache_breakpoint, + }, + ], + }, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_1", + "content": [ + {"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint} + ], + }, + ], + ) + messages_before: Final = copy.deepcopy(messages) + + default_input, default_instructions = handler.convert_chat_completion_messages_to_responses_api(messages) + kept_input, kept_instructions = handler.convert_chat_completion_messages_to_responses_api( + messages, + keep_prompt_cache_breakpoints=True, + ) + + assert default_instructions is None + assert default_input == [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "Review these inputs"}, + {"type": "input_image", "image_url": image_data_url, "detail": "auto"}, + {"type": "input_file", "file_data": file_data, "filename": "input.pdf"}, + ], + }, + { + "type": "function_call", + "call_id": "call_1", + "name": "lookup", + "arguments": "{}", + }, + { + "type": "function_call_output", + "call_id": "call_1", + "output": [{"type": "input_text", "text": "Tool result"}], + }, + ] + assert kept_instructions is None + assert kept_input == [ + { + "type": "message", + "role": "user", + "content": [ + { + "type": "input_text", + "text": "Review these inputs", + "prompt_cache_breakpoint": cache_breakpoint, + }, + { + "type": "input_image", + "image_url": image_data_url, + "detail": "auto", + "prompt_cache_breakpoint": cache_breakpoint, + }, + { + "type": "input_file", + "file_data": file_data, + "filename": "input.pdf", + "prompt_cache_breakpoint": cache_breakpoint, + }, + ], + }, + { + "type": "function_call", + "call_id": "call_1", + "name": "lookup", + "arguments": "{}", + }, + { + "type": "function_call_output", + "call_id": "call_1", + "output": [{"type": "input_text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint}], + }, + ] + assert messages == messages_before + + @pytest.mark.parametrize( ("litellm_params", "keep_marker"), (({"base_model": "gpt-5.6"}, True), ({}, False)), From 77e21ae37a57e4330ff130f9b74f4a49a049e333 Mon Sep 17 00:00:00 2001 From: Yucheng Date: Mon, 5 Oct 2026 00:03:39 -0700 Subject: [PATCH 6/9] fix(responses): read prompt_cache_breakpoint without validating content blocks The marker read ran every chat content block through TypeAdapter(dict[str, object]).validate_python, which rejects dict blocks with non-string keys that chat completion callers passing Python dicts could previously send; the request then failed with a pydantic ValidationError on the bridge keep path, the strip path, and the image/file conversions alike. item is already isinstance-narrowed to a dict at every read site, so read the marker with dict.get directly and drop the adapter. Adds a regression test covering text/image_url/file blocks carrying non-string keys on both the keep (gpt-5.6) and strip (gpt-4o) paths. --- .../transformation.py | 9 ++--- ...responses_transformation_transformation.py | 38 +++++++++++++++++++ 2 files changed, 42 insertions(+), 5 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index d69b2915067..9a710a5ac99 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -19,7 +19,7 @@ from openai.types.responses.response_input_param import ( from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam from openai.types.responses.tool_param import FunctionToolParam -from pydantic import BaseModel, TypeAdapter +from pydantic import BaseModel import litellm from litellm import ModelResponse @@ -78,7 +78,6 @@ _CHAT_COMPLETION_FIELDS: Final = frozenset((*ModelResponse.model_fields, "usage" _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *ResponsesAPIResponse.model_fields)) - frozenset( ChatCompletion.model_fields ) -_CHAT_CONTENT_ITEM: Final = TypeAdapter(dict[str, object]) def _strip_prompt_cache_breakpoint_from_content_block(value: object) -> object: @@ -1118,7 +1117,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if original_type == "text": converted = with_prompt_cache_breakpoint( self._convert_content_str_to_input_text(item.get("text", ""), role), - _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), + item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: text -> %s", converted) @@ -1131,7 +1130,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): role, ) ), - _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), + item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) @@ -1147,7 +1146,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): _input_file_from_file_value( cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked ), - _CHAT_CONTENT_ITEM.validate_python(item).get("prompt_cache_breakpoint"), + item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: file -> %s", converted) 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 fcb44958ffb..ff48ef545cd 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 @@ -4233,6 +4233,44 @@ def test_prompt_cache_breakpoint_survives_chat_to_responses_conversion( assert request["prompt_cache_options"] == cache_breakpoint +def test_prompt_cache_breakpoint_read_tolerates_non_string_content_block_keys() -> None: + handler: Final = LiteLLMResponsesTransformationHandler() + # Non-string keys are not JSON-representable but are accepted by chat completion + # callers passing Python dicts; reading the marker must not validate or reject them. + content: Final = [ + {"type": "text", "text": "Stable prefix", 1: "ignored"}, + {"type": "image_url", "image_url": "https://example.com/image.png", 2: "ignored"}, + {"type": "file", "file": {"file_id": "file-123"}, 3: "ignored"}, + ] + messages: Final = [{"role": "user", "content": content}] + + for model in ("gpt-5.6", "gpt-4o"): # marker keep path and strip path both read the block + request: dict[str, object] = handler.transform_request( + model=model, + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + litellm_logging_obj=Mock(), + ) + + assert request["input"] == [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "Stable prefix"}, + { + "type": "input_image", + "image_url": "https://example.com/image.png", + "detail": "auto", + }, + {"type": "input_file", "file_id": "file-123"}, + ], + } + ] + + def test_prompt_cache_breakpoints_are_dropped_for_unsupported_models() -> None: handler: Final = LiteLLMResponsesTransformationHandler() cache_breakpoint: Final = {"mode": "explicit"} From f6a0299e914b028cb88e214e57a1e5941a0c1723 Mon Sep 17 00:00:00 2001 From: Yucheng Date: Mon, 5 Oct 2026 00:13:24 -0700 Subject: [PATCH 7/9] fix(responses): cast content block before reading prompt_cache_breakpoint basedpyright flags the raw dict.get read as reportUnknownArgumentType (+2 against the error budget); cast the isinstance-narrowed block to dict[str, object] first, matching the strip helpers' cast-ok idiom. --- .../transformation.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 9a710a5ac99..94a71bd6d7d 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1117,7 +1117,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if original_type == "text": converted = with_prompt_cache_breakpoint( self._convert_content_str_to_input_text(item.get("text", ""), role), - item.get("prompt_cache_breakpoint"), + cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + "prompt_cache_breakpoint" + ), ) result.append(converted) verbose_logger.debug("Chat provider: text -> %s", converted) @@ -1130,7 +1132,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): role, ) ), - item.get("prompt_cache_breakpoint"), + cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + "prompt_cache_breakpoint" + ), ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) @@ -1146,7 +1150,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): _input_file_from_file_value( cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked ), - item.get("prompt_cache_breakpoint"), + cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + "prompt_cache_breakpoint" + ), ) result.append(converted) verbose_logger.debug("Chat provider: file -> %s", converted) From 5160b21974833555829c780cb84eac715466e14b Mon Sep 17 00:00:00 2001 From: Yucheng Date: Mon, 5 Oct 2026 00:39:42 -0700 Subject: [PATCH 8/9] style(responses): ruff-format the marker-read cast --- .../transformation.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 94a71bd6d7d..ad25117943b 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1117,7 +1117,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if original_type == "text": converted = with_prompt_cache_breakpoint( self._convert_content_str_to_input_text(item.get("text", ""), role), - cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + cast( + dict[str, object], item + ).get( # cast-ok: isinstance confirms the content block is a mapping "prompt_cache_breakpoint" ), ) @@ -1132,7 +1134,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): role, ) ), - cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + cast( + dict[str, object], item + ).get( # cast-ok: isinstance confirms the content block is a mapping "prompt_cache_breakpoint" ), ) @@ -1150,7 +1154,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): _input_file_from_file_value( cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked ), - cast(dict[str, object], item).get( # cast-ok: isinstance confirms the content block is a mapping + cast( + dict[str, object], item + ).get( # cast-ok: isinstance confirms the content block is a mapping "prompt_cache_breakpoint" ), ) From 90daacfe5a5b8dde2c6d7252c0c9325374bbe45c Mon Sep 17 00:00:00 2001 From: Yucheng Date: Mon, 5 Oct 2026 00:59:44 -0700 Subject: [PATCH 9/9] fix(responses): hoist one cast-ok content block read for the type gates A Final assignment inside the conversion loop trips reportGeneralTypeIssues, and per-site casts trip the LIT006 budget; read the marker through one cast-narrowed local instead. --- .../transformation.py | 21 ++++++------------- 1 file changed, 6 insertions(+), 15 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index ad25117943b..0ed2af1e35c 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1114,14 +1114,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif isinstance(item, dict): # Handle multimodal content original_type = item.get("type") + content_item = cast( # cast-ok: isinstance confirms the content block is a mapping + dict[str, object], item + ) if original_type == "text": converted = with_prompt_cache_breakpoint( self._convert_content_str_to_input_text(item.get("text", ""), role), - cast( - dict[str, object], item - ).get( # cast-ok: isinstance confirms the content block is a mapping - "prompt_cache_breakpoint" - ), + content_item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: text -> %s", converted) @@ -1134,11 +1133,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): role, ) ), - cast( - dict[str, object], item - ).get( # cast-ok: isinstance confirms the content block is a mapping - "prompt_cache_breakpoint" - ), + content_item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) @@ -1154,11 +1149,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): _input_file_from_file_value( cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked ), - cast( - dict[str, object], item - ).get( # cast-ok: isinstance confirms the content block is a mapping - "prompt_cache_breakpoint" - ), + content_item.get("prompt_cache_breakpoint"), ) result.append(converted) verbose_logger.debug("Chat provider: file -> %s", converted)