fix(responses): keep prompt_cache_breakpoint markers in the chat to responses bridge (#44119)

* 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 <simonsorg13@gmail.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(responses): honor base_model when gating bridge cache breakpoints

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(responses): avoid recursive cache-breakpoint stripping

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(responses): justify bridge stripping casts

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* 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>

* 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.

* 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.

* style(responses): ruff-format the marker-read cast

* 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.

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Simon Sorg <simonsorg13@gmail.com>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 16:25:36 -07:00 • committed by GitHub
parent 225fd3bcdf
commit d910653b33
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 1290 additions and 11 deletions

View file

@ -24,6 +24,7 @@ from pydantic import BaseModel
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.hidden_params import get_hidden_params, get_or_create_hidden_params
from litellm.litellm_core_utils.prompt_templates.common_utils import (
responses_reasoning_items_from_thinking_blocks,
@ -80,6 +81,38 @@ _RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *Response
)
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) # 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) # 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) # 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)
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) # 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()
if key != "prompt_cache_breakpoint"
}
def _strip_prompt_cache_breakpoints(input_items: list[object]) -> list[object]:
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]:
return MappingProxyType(
{
@ -364,6 +397,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]] = []
@ -594,24 +641,31 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
litellm_logging_obj: "LiteLLMLoggingObj",
client: object | None = None,
) -> dict:
(
input_items,
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)
)
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.
if not input_items and instructions is not None:
input_items = [
is_system_only_request: Final = not converted_input_items 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 converted_input_items
)
instructions: Final = None if is_system_only_request else converted_instructions
optional_params = self._extract_extra_body_params(optional_params)

View file

@ -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

View file

@ -0,0 +1,201 @@
from __future__ import annotations
import uuid
from typing import Final
import httpx
import pytest
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))
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

View file

@ -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()

View file

@ -1,8 +1,9 @@
import copy
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
@ -21,7 +22,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
@ -4404,6 +4405,295 @@ 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"}
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_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"}}],
}
]
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)),
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()