fix(responses): keep cache breakpoints on blocks the bridge stringifies (#40032)

* fix(responses): keep cache breakpoints on requests that reach the Responses API

The chat-completions bridge rebuilt every content block without the
prompt_cache_breakpoint marker the cache-control hook had just placed on it, so
the request went out with prompt_cache_options set to explicit mode and nothing
actually marked. Explicit mode caches only what is marked, so those deployments
lost the implicit caching they were getting before the injection point was added

The native surface had a second miss. The bridging check looked the provider
config up with the still-prefixed model while the config keys on the bare one,
so every bedrock_mantle model read as having no native Responses support, and
role-targeted injection points were deferred to a chat-completions pass that
never runs for them

* fix(responses): keep prompt cache breakpoint on stringified content blocks

* test: cover the responses/ routing prefix in the native cache point case

* fix(responses): keep only the fallback-branch marker carry, main already routes the hook

#42281 landed the text, image_url and file branch carry and the implicit default
on main, and the Responses routing lookup this branch changed has no observable
effect there, so the PR shrinks to the unknown-block fallback branch and its
regression test

* fix(responses): drop a malformed prompt_cache_breakpoint under drop_params in the chat-to-Responses bridge

* fix(responses): accept the 30m prompt_cache_breakpoint ttl under drop_params

OpenAI's Responses API takes a ttl of 30m on an explicit breakpoint marker, so the drop_params validation keeps it instead of stripping it as an unknown key. Trims the bridge test docstrings to the dated vendor citation.

* test(integration): audit cells for prompt_cache_breakpoint through the chat-to-Responses bridge

Adds the /audit cells for the bridge's prompt_cache_breakpoint carry: valid and malformed markers on every
carrying block kind and on tool and assistant content, the drop_params on and off contract, the hook-injected
system marker, the Anthropic SDK path through the Responses adapter, upstream errors, idempotent spend rows,
and the chaos cells (mixed burst, upstream outage, slow streams, worker SIGKILL and a proxy restart mid burst),
all against a scripted Responses endpoint with the request it received read back by response id

* test(integration): give every scripted responses reply its own id

The prompt-cache-breakpoint audit's scripted upstream minted the response id
from the request's marker, so three identical marked requests shared one id.
The spend-log writer skips rows whose request_id already landed, and the
idempotent-logging cell saw one row for three calls on every leg. The upstream
now mints a unique id per reply like a real provider, and the cells match a
response to its request by the marker inside that id.
This commit is contained in:
Mateo Wang 2026-10-07 18:28:39 -07:00 • committed by GitHub
parent c52082e10c
commit 9d3bf29d6d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1383 additions and 297 deletions

View file

@ -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
from pydantic import BaseModel, TypeAdapter, ValidationError
import litellm
from litellm import ModelResponse
@ -46,6 +46,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParamFunctionChunk,
PromptCacheBreakpoint,
Reasoning,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
@ -238,6 +239,19 @@ def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Li
return "length"
_PROMPT_CACHE_BREAKPOINT: Final = TypeAdapter(PromptCacheBreakpoint)
def _prompt_cache_breakpoint_for_wire(marker: object, drop_params: bool) -> object:
if marker is None or not drop_params:
return marker
try:
return _PROMPT_CACHE_BREAKPOINT.validate_python(marker)
except ValidationError:
verbose_logger.debug("Chat provider: dropping malformed prompt_cache_breakpoint %r under drop_params", marker)
return None
def _input_file_from_file_value(file_value: object) -> dict[str, object]:
if not isinstance(file_value, dict):
return {"type": "input_file"}
@ -400,9 +414,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
self,
messages: list["AllMessageValues"],
*,
drop_params: bool = False,
keep_prompt_cache_breakpoints: bool = False,
) -> tuple[list[object], str | None]:
converted_input_items, instructions = self._convert_chat_completion_messages_to_responses_input(messages)
converted_input_items, instructions = self._convert_chat_completion_messages_to_responses_input(
messages, drop_params=drop_params
)
return (
converted_input_items
if keep_prompt_cache_breakpoints
@ -411,7 +428,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
)
def _convert_chat_completion_messages_to_responses_input(
self, messages: list["AllMessageValues"]
self, messages: list["AllMessageValues"], *, drop_params: bool = False
) -> tuple[list[object], str | None]:
input_items: Final[list[object]] = []
instructions: str | None = None
@ -452,6 +469,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
"content": self._convert_content_to_responses_format(
content,
role,
drop_params=drop_params,
),
}
)
@ -470,6 +488,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
tool_output = self._convert_content_to_responses_format(
content,
"user", # Use "user" role to get input_* types
drop_params=drop_params,
)
else:
# Fallback: convert unexpected types to input_text
@ -497,7 +516,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
{
"type": "message",
"role": "assistant",
"content": self._convert_content_to_responses_format(content, "assistant"),
"content": self._convert_content_to_responses_format(
content, "assistant", drop_params=drop_params
),
}
)
for tool_call in tool_calls:
@ -531,7 +552,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
{
"type": "message",
"role": role,
"content": self._convert_content_to_responses_format(content, cast(str, role)),
"content": self._convert_content_to_responses_format(
content, cast(str, role), drop_params=drop_params
),
}
)
elif role == "assistant":
@ -647,6 +670,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
)
converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(
messages,
drop_params=bool(litellm_params.get("drop_params") or litellm.drop_params),
keep_prompt_cache_breakpoints=supports_prompt_cache_breakpoint,
)
# OpenAI's Responses API rejects an empty input. For a system-only
@ -1126,6 +1150,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
]
| None,
role: str,
drop_params: bool = False,
) -> list[dict[str, object]]:
"""Convert chat completion content to responses API format"""
from litellm.types.llms.openai import ChatCompletionImageObject
@ -1152,7 +1177,7 @@ 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"),
_prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params),
)
result.append(converted)
verbose_logger.debug("Chat provider: text -> %s", converted)
@ -1165,7 +1190,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
cast(ChatCompletionImageObject, item), role
),
),
item.get("prompt_cache_breakpoint"),
_prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params),
)
result.append(converted)
verbose_logger.debug("Chat provider: image_url -> %s", converted)
@ -1181,7 +1206,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
_input_file_from_file_value(
cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked
),
item.get("prompt_cache_breakpoint"),
_prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params),
)
result.append(converted)
verbose_logger.debug("Chat provider: file -> %s", converted)
@ -1203,7 +1228,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
verbose_logger.debug("Chat provider: passthrough -> %s", item)
else:
# Default to input_text for unknown types
converted = self._convert_content_str_to_input_text(str(item.get("text", item)), role)
converted = with_prompt_cache_breakpoint(
self._convert_content_str_to_input_text(str(item.get("text", item)), role),
_prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params),
)
result.append(converted)
verbose_logger.debug("Chat provider: unknown(%s) -> %s", original_type, converted)
verbose_logger.debug("Chat provider: Final converted content: %s", result)

View file

@ -652,6 +652,7 @@ class ChatCompletionCachedContent(TypedDict):
class PromptCacheBreakpoint(TypedDict):
mode: ReadOnly[Literal["explicit"]]
ttl: NotRequired[ReadOnly[Literal["30m"]]]
class PromptCacheOptions(TypedDict, total=False):

View file

@ -0,0 +1,210 @@
from __future__ import annotations
import os
import re
import uuid
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Literal, TypeAlias, assert_never
from urllib.parse import urlsplit
import psutil
import psycopg
from integration._support import responses_vendor as rv
from integration._support.client import eventually, object_value, string_value
from integration._support.database import ROWS
from integration._support.openai_wire import answering_model_discovery, responses_reply
from integration._support.wire import Reply, Request, Wire
from psycopg.rows import DictRow, dict_row
from pydantic import JsonValue
MODEL: Final = "openai/responses/gpt-6.1-sol"
EXPLICIT: Final[Mapping[str, JsonValue]] = {"mode": "explicit"}
EXPLICIT_30M: Final[Mapping[str, JsonValue]] = {"mode": "explicit", "ttl": "30m"}
NO_CACHE: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}}
INJECTION: Final[Mapping[str, JsonValue]] = {
"cache_control_injection_points": [{"location": "message", "role": "system"}],
"prompt_cache_options": {"mode": "explicit"},
}
_SCRIPTED_FAILURE: Final = re.compile(r"fail-(\d{3})")
_MINTED_RESPONSE: Final = re.compile(r"^resp_([0-9a-f]{32})-[0-9a-f]{32}$")
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
Kind: TypeAlias = Literal["text", "image_url", "file", "input_audio"]
KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "input_audio")
WIRE_TYPE: Final[Mapping[Kind, str]] = {
"text": "input_text",
"image_url": "input_image",
"file": "input_file",
"input_audio": "input_text",
}
def _scripted(request: Request) -> Reply:
body: Final = rv.JSON_OBJECT.validate_json(request.body)
text: Final = request.body.decode()
marker: Final = rv.newest_marker(text)
failure: Final = _SCRIPTED_FAILURE.search(text)
if failure is not None:
return rv.error(int(failure.group(1)), f"scripted {failure.group(1)} marker-{marker}", "scripted_failure")
return responses_reply(
f"resp_{marker or uuid.uuid4().hex}-{uuid.uuid4().hex}",
string_value(body["model"]),
rv.answer(marker),
stream=body.get("stream") is True,
)
respond: Final = answering_model_discovery(_scripted)
def response_marker(identity: str) -> str | None:
minted: Final = tuple(
found for candidate in rv.response_identities(identity) if (found := _MINTED_RESPONSE.match(candidate))
)
return minted[0].group(1) if minted else None
def answers(identity: str, marker: str) -> bool:
return response_marker(identity) == marker
def prompt(marker: str) -> str:
return f"Say marker-{marker}"
def text(value: str) -> dict[str, JsonValue]:
return {"type": "text", "text": value}
def marked(block: Mapping[str, JsonValue], marker: JsonValue) -> dict[str, JsonValue]:
return {**block, "prompt_cache_breakpoint": marker}
def block(kind: Kind, value: str) -> dict[str, JsonValue]:
match kind:
case "text":
return text(value)
case "image_url":
return {"type": "image_url", "image_url": {"url": "https://example.com/breakpoint.png"}}
case "file":
return {"type": "file", "file": {"file_id": "file-breakpoint"}}
case "input_audio":
return {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}
case _:
assert_never(kind)
def drained_posts(wire: Wire) -> tuple[Request, ...]:
return tuple(request for request in wire.drain() if request.method == "POST")
def with_marker(posts: Sequence[Request], marker: str) -> tuple[Request, ...]:
return tuple(request for request in posts if f"marker-{marker}" in request.body.decode())
def posted(wire: Wire, marker: str) -> Request:
matching: Final = with_marker(drained_posts(wire), marker)
assert len(matching) == 1, [request.body for request in matching]
(request,) = matching
assert request.target == "/v1/responses", request.target
return request
def body_of(request: Request) -> dict[str, JsonValue]:
return rv.JSON_OBJECT.validate_json(request.body)
def input_items(request: Request) -> list[dict[str, JsonValue]]:
return rv.ITEMS.validate_python(body_of(request)["input"])
def content_of(items: Sequence[Mapping[str, JsonValue]], role: str) -> list[dict[str, JsonValue]]:
messages: Final = tuple(item for item in items if item.get("type") == "message" and item.get("role") == role)
assert len(messages) == 1, items
return rv.ITEMS.validate_python(messages[0]["content"])
def single_block(items: Sequence[Mapping[str, JsonValue]], role: str) -> dict[str, JsonValue]:
blocks: Final = content_of(items, role)
assert len(blocks) == 1, blocks
return blocks[0]
def instruction_block(items: Sequence[Mapping[str, JsonValue]]) -> dict[str, JsonValue]:
messages: Final = tuple(
item for item in items if item.get("type") == "message" and item.get("role") in ("system", "developer")
)
assert len(messages) == 1, items
blocks: Final = rv.ITEMS.validate_python(messages[0]["content"])
assert len(blocks) == 1, blocks
return blocks[0]
def function_output(items: Sequence[Mapping[str, JsonValue]], call_id: str) -> list[dict[str, JsonValue]]:
outputs: Final = tuple(
item for item in items if item.get("type") == "function_call_output" and item.get("call_id") == call_id
)
assert len(outputs) == 1, items
return rv.ITEMS.validate_python(outputs[0]["output"])
def assert_marker(block_on_wire: Mapping[str, JsonValue], expected: JsonValue) -> None:
if expected is None:
assert "prompt_cache_breakpoint" not in block_on_wire, block_on_wire
return
assert block_on_wire.get("prompt_cache_breakpoint") == expected, block_on_wire
@dataclass(frozen=True, slots=True)
class SpendLogs:
connection: psycopg.Connection[DictRow]
def rows_for(self, model: str) -> list[dict[str, JsonValue]]:
cursor: Final = self.connection.execute(
'SELECT litellm_call_id, request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group = %s', (model,)
)
return ROWS.validate_python(cursor.fetchall())
def landed(
self, model: str, call_id: str, marker: str | None, *, status: str = "success", seconds: float = 70
) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: self.rows_for(model),
lambda found: any(row["litellm_call_id"] == call_id for row in found),
seconds=seconds,
)
matching: Final = tuple(row for row in rows if row["litellm_call_id"] == call_id)
assert len(matching) == 1, rows
(row,) = matching
assert row["status"] == status, row
assert marker is None or answers(string_value(row["request_id"]), marker), (row, marker)
return row
@contextmanager
def spend_logs() -> Iterator[SpendLogs]:
with psycopg.connect(os.environ["DATABASE_URL"], row_factory=dict_row, autocommit=True) as connection:
connection.execute("SET default_transaction_read_only = on")
yield SpendLogs(connection)
def model_id(entries: Sequence[JsonValue], model: str) -> str:
matching: Final = tuple(entry for entry in entries if object_value(entry).get("model_name") == model)
assert len(matching) == 1, entries
return string_value(object_value(object_value(matching[0])["model_info"])["id"])
def started_worker_pids(log: Path) -> tuple[int, ...]:
return tuple(int(found.group(1)) for found in _STARTED_WORKER.finditer(log.read_text()))
def open_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
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
)

View file

@ -0,0 +1,510 @@
import asyncio
import uuid
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Final, Literal, TypeAlias
import anthropic
import httpx
import openai
import pytest
from integration._support import prompt_cache_breakpoint as pcb
from integration._support import responses_vendor as rv
from integration._support.client import Gateway, eventually, gateway_from_environment, string_value
from integration._support.wire import Request, Wire, wire_server
from pydantic import JsonValue
pytestmark: Final = pytest.mark.timeout(120)
Mode: TypeAlias = Literal["on", "off"]
_UNKNOWN_KEY: Final[Mapping[str, JsonValue]] = {"mode": "explicit", "note": "kept"}
_MALFORMED: Final[tuple[tuple[str, JsonValue], ...]] = (
("string", "yes"),
("int", 1),
("list", ["explicit"]),
("empty-string", ""),
("5kb-string", "x" * 5000),
("empty-object", {}),
("bogus-mode", {"mode": "bogus"}),
("bad-ttl", {"mode": "explicit", "ttl": "1h"}),
)
_MALFORMED_IDS: Final = tuple(name for name, _ in _MALFORMED)
_MALFORMED_VALUES: Final = tuple(value for _, value in _MALFORMED)
_CASES: Final[tuple[tuple[str, Mode, JsonValue, JsonValue], ...]] = (
("valid-on", "on", pcb.EXPLICIT, pcb.EXPLICIT),
("valid-off", "off", pcb.EXPLICIT, pcb.EXPLICIT),
("malformed-on", "on", "yes", None),
("malformed-off", "off", "yes", "yes"),
)
_CASE_IDS: Final = tuple(case[0] for case in _CASES)
_CASE_VALUES: Final = tuple(case[1:] for case in _CASES)
_ADAPTER_CASES: Final[tuple[tuple[str, Mode, JsonValue, JsonValue], ...]] = (
("valid-on", "on", pcb.EXPLICIT, pcb.EXPLICIT),
("valid-off", "off", pcb.EXPLICIT, pcb.EXPLICIT),
("malformed-on", "on", "yes", "yes"),
("malformed-off", "off", "yes", "yes"),
)
_ADAPTER_CASE_IDS: Final = tuple(case[0] for case in _ADAPTER_CASES)
_ADAPTER_CASE_VALUES: Final = tuple(case[1:] for case in _ADAPTER_CASES)
@dataclass(frozen=True, slots=True)
class _Bridge:
gateway: Gateway
wire: Wire
on: str
off: str
injecting_on: str
injecting_off: str
spend: pcb.SpendLogs
def model(self, mode: Mode) -> str:
return self.on if mode == "on" else self.off
def injecting(self, mode: Mode) -> str:
return self.injecting_on if mode == "on" else self.injecting_off
@property
def api_base(self) -> str:
return f"{self.wire.url}/v1"
@pytest.fixture(scope="module")
def bridge() -> Iterator[_Bridge]:
with (
wire_server(pcb.respond) as wire,
gateway_from_environment() as gateway,
gateway.scenario() as scenario,
pcb.spend_logs() as spend,
):
api_base: Final = f"{wire.url}/v1"
yield _Bridge(
gateway,
wire,
scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True),
scenario.model(model=pcb.MODEL, api_base=api_base),
scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True, **pcb.INJECTION),
scenario.model(model=pcb.MODEL, api_base=api_base, **pcb.INJECTION),
spend,
)
def _v1(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/") + "/v1"
def _user(marker: str, breakpoint: JsonValue) -> dict[str, JsonValue]:
return {"role": "user", "content": [pcb.marked(pcb.text(pcb.prompt(marker)), breakpoint)]}
def _chat(
bridge: _Bridge, model: str, messages: Sequence[JsonValue], *, stream: bool = False, key: str | None = None
) -> httpx.Response:
return bridge.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": list(messages), "stream": stream, **pcb.NO_CACHE},
key=key,
)
def _completion(response: httpx.Response, marker: str) -> str:
assert response.status_code == 200, response.text
body: Final = rv.JSON_OBJECT.validate_json(response.text)
assert pcb.answers(string_value(body["id"]), marker), body
(choice,) = rv.ITEMS.validate_python(body["choices"])
assert rv.JSON_OBJECT.validate_python(choice["message"])["content"] == rv.answer(marker), body
return response.headers["x-litellm-call-id"]
def _wire_body(request: Request, *, stream: bool = False) -> dict[str, JsonValue]:
body: Final = pcb.body_of(request)
assert body["model"] == "gpt-6.1-sol", body
assert (body.get("stream") is True) is stream, body
return body
def _user_block_on_wire(bridge: _Bridge, marker: str, *, stream: bool = False) -> dict[str, JsonValue]:
request: Final = pcb.posted(bridge.wire, marker)
block: Final = pcb.single_block(pcb.input_items(request), "user")
_wire_body(request, stream=stream)
assert block["type"] == "input_text" and block["text"] == pcb.prompt(marker), block
return block
def test_openai_sdk_sends_a_valid_marker_through_the_bridge(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
with openai.OpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client:
raw: Final = client.chat.completions.with_raw_response.create(
model=bridge.on, messages=[_user(marker, pcb.EXPLICIT)], extra_body=dict(pcb.NO_CACHE)
)
completion: Final = raw.parse()
assert pcb.answers(completion.id, marker), completion
assert completion.choices[0].message.content == rv.answer(marker), completion
pcb.assert_marker(_user_block_on_wire(bridge, marker), pcb.EXPLICIT)
bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], marker)
def test_openai_sdk_stream_carries_the_system_list_marker(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.marked(pcb.text("sys"), pcb.EXPLICIT)]}
with openai.OpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client:
raw: Final = client.chat.completions.with_raw_response.create(
model=bridge.on,
messages=[system, {"role": "user", "content": pcb.prompt(marker)}],
stream=True,
extra_body=dict(pcb.NO_CACHE),
)
chunks: Final = tuple(raw.parse())
assert chunks and pcb.answers(chunks[0].id, marker), chunks
streamed: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
assert streamed == rv.answer(marker), chunks
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request, stream=True)
(system_block,) = pcb.content_of(pcb.input_items(request), "system")
assert system_block["type"] == "input_text" and system_block["text"] == "sys", system_block
pcb.assert_marker(system_block, pcb.EXPLICIT)
bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], marker)
async def test_async_openai_sdk_keeps_the_ttl_without_drop_params(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
async with openai.AsyncOpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client:
raw: Final = await client.chat.completions.with_raw_response.create(
model=bridge.off, messages=[_user(marker, pcb.EXPLICIT_30M)], extra_body=dict(pcb.NO_CACHE)
)
completion: Final = raw.parse()
assert pcb.answers(completion.id, marker), completion
assert completion.choices[0].message.content == rv.answer(marker), completion
pcb.assert_marker(_user_block_on_wire(bridge, marker), pcb.EXPLICIT_30M)
bridge.spend.landed(bridge.off, raw.headers["x-litellm-call-id"], marker)
@pytest.mark.parametrize("mode", ("on", "off"))
@pytest.mark.parametrize("breakpoint", (pcb.EXPLICIT, pcb.EXPLICIT_30M), ids=("explicit", "ttl"))
def test_valid_marker_shapes_reach_the_wire_unchanged(bridge: _Bridge, mode: Mode, breakpoint: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, breakpoint)]), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), breakpoint)
bridge.spend.landed(bridge.model(mode), call_id, marker)
@pytest.mark.parametrize(
("mode", "expected"), (("on", pcb.EXPLICIT), ("off", _UNKNOWN_KEY)), ids=("normalized-on", "verbatim-off")
)
def test_marker_with_an_unknown_key(bridge: _Bridge, mode: Mode, expected: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, _UNKNOWN_KEY)]), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), expected)
bridge.spend.landed(bridge.model(mode), call_id, marker)
def _second_block_on_wire(bridge: _Bridge, marker: str) -> dict[str, JsonValue]:
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request)
first, second = pcb.content_of(pcb.input_items(request), "user")
assert first == {"type": "input_text", "text": pcb.prompt(marker)}, first
return second
@pytest.mark.parametrize("mode", ("on", "off"))
@pytest.mark.parametrize("kind", pcb.KINDS)
def test_valid_marker_is_carried_on_every_block_kind(bridge: _Bridge, kind: pcb.Kind, mode: Mode) -> None:
marker: Final = uuid.uuid4().hex
content: Final[list[JsonValue]] = [
pcb.text(pcb.prompt(marker)),
pcb.marked(pcb.block(kind, "second"), pcb.EXPLICIT),
]
call_id: Final = _completion(_chat(bridge, bridge.model(mode), [{"role": "user", "content": content}]), marker)
second: Final = _second_block_on_wire(bridge, marker)
assert second["type"] == pcb.WIRE_TYPE[kind], second
pcb.assert_marker(second, pcb.EXPLICIT)
bridge.spend.landed(bridge.model(mode), call_id, marker)
@pytest.mark.parametrize("kind", pcb.KINDS)
def test_malformed_marker_is_dropped_on_every_block_kind(bridge: _Bridge, kind: pcb.Kind) -> None:
marker: Final = uuid.uuid4().hex
content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.block(kind, "second"), "yes")]
call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker)
second: Final = _second_block_on_wire(bridge, marker)
assert second["type"] == pcb.WIRE_TYPE[kind], second
pcb.assert_marker(second, None)
bridge.spend.landed(bridge.on, call_id, marker)
@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _CASE_VALUES, ids=_CASE_IDS)
def test_tool_output_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
messages: Final[list[JsonValue]] = [
{"role": "user", "content": pcb.prompt(marker)},
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "call_1", "content": [pcb.marked(pcb.text("found it"), breakpoint)]},
]
call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker)
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request)
(output,) = pcb.function_output(pcb.input_items(request), "call_1")
assert output["type"] == "input_text" and output["text"] == "found it", output
pcb.assert_marker(output, expected)
bridge.spend.landed(bridge.model(mode), call_id, marker)
@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _CASE_VALUES, ids=_CASE_IDS)
def test_assistant_list_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
messages: Final[list[JsonValue]] = [
{"role": "user", "content": pcb.prompt(marker)},
{"role": "assistant", "content": [pcb.marked(pcb.text("earlier answer"), breakpoint)]},
{"role": "user", "content": "and again"},
]
call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker)
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request)
(earlier,) = pcb.content_of(pcb.input_items(request), "assistant")
assert earlier["type"] == "output_text" and earlier["text"] == "earlier answer", earlier
pcb.assert_marker(earlier, expected)
bridge.spend.landed(bridge.model(mode), call_id, marker)
def test_injected_system_marker_survives_a_trailing_audio_block(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.block("input_audio", "")]}
messages: Final[list[JsonValue]] = [system, {"role": "user", "content": pcb.prompt(marker)}]
call_id: Final = _completion(_chat(bridge, bridge.injecting_on, messages), marker)
request: Final = pcb.posted(bridge.wire, marker)
body: Final = _wire_body(request)
assert body["prompt_cache_options"] == {"mode": "explicit"}, body
first, audio = pcb.content_of(pcb.input_items(request), "system")
assert first == {"type": "input_text", "text": "sys"}, first
assert audio["type"] == "input_text", audio
assert string_value(audio["text"]).startswith("{'type': 'input_audio'"), audio
pcb.assert_marker(audio, pcb.EXPLICIT)
bridge.spend.landed(bridge.injecting_on, call_id, marker)
def test_injected_marker_on_a_string_system_message(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
messages: Final[list[JsonValue]] = [
{"role": "system", "content": "Answer briefly"},
{"role": "user", "content": pcb.prompt(marker)},
]
call_id: Final = _completion(_chat(bridge, bridge.injecting_off, messages), marker)
request: Final = pcb.posted(bridge.wire, marker)
body: Final = _wire_body(request)
assert body["prompt_cache_options"] == {"mode": "explicit"}, body
system_block: Final = pcb.single_block(pcb.input_items(request), "system")
assert system_block == {"type": "input_text", "text": "Answer briefly", "prompt_cache_breakpoint": pcb.EXPLICIT}
bridge.spend.landed(bridge.injecting_off, call_id, marker)
def _anthropic(bridge: _Bridge) -> anthropic.Anthropic:
return anthropic.Anthropic(base_url=str(bridge.gateway.client.base_url), api_key=bridge.gateway.key, max_retries=0)
def _anthropic_text(message: anthropic.types.Message, marker: str) -> None:
(content,) = message.content
assert content.type == "text" and content.text == rv.answer(marker), message
@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _ADAPTER_CASE_VALUES, ids=_ADAPTER_CASE_IDS)
def test_anthropic_sdk_marker_on_user_text(
bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue
) -> None:
marker: Final = uuid.uuid4().hex
with _anthropic(bridge) as client:
raw: Final = client.messages.with_raw_response.create(
model=bridge.model(mode),
max_tokens=64,
messages=[{"role": "user", "content": [pcb.marked(pcb.text(pcb.prompt(marker)), breakpoint)]}],
extra_body=dict(pcb.NO_CACHE),
)
_anthropic_text(raw.parse(), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), expected)
bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], None)
def test_anthropic_sdk_system_string_gets_the_injected_marker(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
with _anthropic(bridge) as client:
raw: Final = client.messages.with_raw_response.create(
model=bridge.injecting_off,
max_tokens=64,
system="Answer briefly",
messages=[{"role": "user", "content": pcb.prompt(marker)}],
extra_body=dict(pcb.NO_CACHE),
)
_anthropic_text(raw.parse(), marker)
request: Final = pcb.posted(bridge.wire, marker)
body: Final = _wire_body(request)
assert body["prompt_cache_options"] == {"mode": "explicit"}, body
instruction: Final = pcb.instruction_block(pcb.input_items(request))
assert instruction == {"type": "input_text", "text": "Answer briefly", "prompt_cache_breakpoint": pcb.EXPLICIT}
bridge.spend.landed(bridge.injecting_off, raw.headers["x-litellm-call-id"], None)
def test_native_responses_request_never_enters_the_bridge(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
block: Final[dict[str, JsonValue]] = {"type": "input_text", "text": pcb.prompt(marker)}
response: Final = bridge.gateway.request(
"POST",
"/v1/responses",
{
"model": bridge.on,
"input": [{"type": "message", "role": "user", "content": [pcb.marked(block, _UNKNOWN_KEY)]}],
**pcb.NO_CACHE,
},
)
assert response.status_code == 200, response.text
body: Final = rv.JSON_OBJECT.validate_json(response.text)
assert pcb.answers(string_value(body["id"]), marker), body
assert rv.answer(marker) in response.text, response.text
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request)
on_wire: Final = pcb.single_block(pcb.input_items(request), "user")
assert on_wire == pcb.marked(block, _UNKNOWN_KEY), on_wire
bridge.spend.landed(bridge.on, response.headers["x-litellm-call-id"], marker)
@pytest.mark.parametrize("breakpoint", _MALFORMED_VALUES, ids=_MALFORMED_IDS)
def test_malformed_marker_is_dropped_under_drop_params(bridge: _Bridge, breakpoint: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, breakpoint)]), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), None)
bridge.spend.landed(bridge.on, call_id, marker)
@pytest.mark.parametrize("breakpoint", _MALFORMED_VALUES, ids=_MALFORMED_IDS)
def test_malformed_marker_passes_verbatim_without_drop_params(bridge: _Bridge, breakpoint: JsonValue) -> None:
marker: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, breakpoint)]), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), breakpoint)
bridge.spend.landed(bridge.off, call_id, marker)
def test_two_marked_blocks_are_both_carried(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
content: Final[list[JsonValue]] = [
pcb.marked(pcb.text(pcb.prompt(marker)), pcb.EXPLICIT),
pcb.marked(pcb.text("and more"), pcb.EXPLICIT_30M),
]
call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker)
request: Final = pcb.posted(bridge.wire, marker)
_wire_body(request)
first, second = pcb.content_of(pcb.input_items(request), "user")
assert first == {"type": "input_text", "text": pcb.prompt(marker), "prompt_cache_breakpoint": pcb.EXPLICIT}
assert second == {"type": "input_text", "text": "and more", "prompt_cache_breakpoint": pcb.EXPLICIT_30M}
bridge.spend.landed(bridge.on, call_id, marker)
def test_wrong_key_is_refused_before_the_wire(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
response: Final = _chat(bridge, bridge.on, [_user(marker, pcb.EXPLICIT)], key="sk-wrong")
assert response.status_code == 401, response.text
assert pcb.with_marker(pcb.drained_posts(bridge.wire), marker) == (), marker
@pytest.mark.parametrize(
("mode", "status"), (("on", 400), ("off", 400), ("on", 401)), ids=("400-on", "400-off", "401-on")
)
def test_upstream_error_reaches_the_caller_once(bridge: _Bridge, mode: Mode, status: int) -> None:
marker: Final = uuid.uuid4().hex
failing: Final[dict[str, JsonValue]] = {
"role": "user",
"content": [pcb.marked(pcb.text(f"{pcb.prompt(marker)} fail-{status}"), pcb.EXPLICIT)],
}
response: Final = _chat(bridge, bridge.model(mode), [failing])
assert response.status_code == status, response.text
assert f"scripted {status} marker-{marker}" in response.text, response.text
(request,) = pcb.with_marker(pcb.drained_posts(bridge.wire), marker)
pcb.assert_marker(pcb.single_block(pcb.input_items(request), "user"), pcb.EXPLICIT)
bridge.spend.landed(bridge.model(mode), response.headers["x-litellm-call-id"], None, status="failure")
follow_up: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(follow_up, pcb.EXPLICIT)]), follow_up)
pcb.assert_marker(_user_block_on_wire(bridge, follow_up), pcb.EXPLICIT)
bridge.spend.landed(bridge.model(mode), call_id, follow_up)
def test_null_drop_params_on_the_deployment_means_off(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
with bridge.gateway.scenario() as scenario:
model: Final = scenario.model(model=pcb.MODEL, api_base=bridge.api_base, drop_params=None)
call_id: Final = _completion(_chat(bridge, model, [_user(marker, "yes")]), marker)
pcb.assert_marker(_user_block_on_wire(bridge, marker), "yes")
bridge.spend.landed(model, call_id, marker)
@pytest.mark.parametrize("mode", ("on", "off"))
@pytest.mark.parametrize("shape", ("null", "missing"))
def test_null_or_missing_marker_sends_a_plain_block(bridge: _Bridge, shape: str, mode: Mode) -> None:
marker: Final = uuid.uuid4().hex
block: Final = pcb.marked(pcb.text(pcb.prompt(marker)), None) if shape == "null" else pcb.text(pcb.prompt(marker))
call_id: Final = _completion(_chat(bridge, bridge.model(mode), [{"role": "user", "content": [block]}]), marker)
assert _user_block_on_wire(bridge, marker) == {"type": "input_text", "text": pcb.prompt(marker)}
bridge.spend.landed(bridge.model(mode), call_id, marker)
async def _send_marked(client: httpx.AsyncClient, key: str, model: str, marker: str) -> httpx.Response:
return await client.post(
"/v1/chat/completions",
json={"model": model, "messages": [_user(marker, "yes")], **pcb.NO_CACHE},
headers={"Authorization": f"Bearer {key}"},
)
def _probe_marker(bridge: _Bridge, model: str) -> JsonValue:
marker: Final = uuid.uuid4().hex
_completion(_chat(bridge, model, [_user(marker, "yes")]), marker)
return _user_block_on_wire(bridge, marker).get("prompt_cache_breakpoint")
@pytest.mark.timeout(180)
async def test_flipping_drop_params_mid_burst_keeps_every_marked_request_answered(bridge: _Bridge) -> None:
gateway: Final = bridge.gateway
with gateway.scenario() as scenario:
model: Final = scenario.model(model=pcb.MODEL, api_base=bridge.api_base, drop_params=True)
identity: Final = pcb.model_id(gateway.get("/model/info")["data"], model)
markers: Final = tuple(uuid.uuid4().hex for _ in range(20))
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
burst: Final = asyncio.gather(*(_send_marked(client, gateway.key, model, marker) for marker in markers))
updated: Final = await asyncio.to_thread(
gateway.request,
"POST",
"/model/update",
{
"model_name": model,
"litellm_params": {"model": pcb.MODEL, "drop_params": False},
"model_info": {"id": identity},
},
)
responses: Final = await burst
assert updated.status_code == 200, updated.text
for marker, response in zip(markers, responses, strict=True):
_completion(response, marker)
posts: Final = pcb.drained_posts(bridge.wire)
for marker in markers:
(request,) = pcb.with_marker(posts, marker)
seen: Final = pcb.single_block(pcb.input_items(request), "user").get("prompt_cache_breakpoint")
assert seen in (None, "yes"), request.body
flipped: Final = eventually(lambda: _probe_marker(bridge, model), lambda seen: seen == "yes", seconds=70)
assert flipped == "yes"
for marker, response in zip(markers, responses, strict=True):
bridge.spend.landed(model, response.headers["x-litellm-call-id"], marker)
def test_three_identical_marked_requests_are_each_sent_and_logged(bridge: _Bridge) -> None:
marker: Final = uuid.uuid4().hex
responses: Final = tuple(_chat(bridge, bridge.on, [_user(marker, pcb.EXPLICIT)]) for _ in range(3))
call_ids: Final = tuple(_completion(response, marker) for response in responses)
assert len(set(call_ids)) == 3, call_ids
posts: Final = pcb.with_marker(pcb.drained_posts(bridge.wire), marker)
assert len(posts) == 3, [request.body for request in posts]
for request in posts:
pcb.assert_marker(pcb.single_block(pcb.input_items(request), "user"), pcb.EXPLICIT)
for call_id in call_ids:
bridge.spend.landed(bridge.on, call_id, marker)

View file

@ -0,0 +1,397 @@
import asyncio
import dataclasses
import signal
import threading
import uuid
from collections import Counter
from collections.abc import Callable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
import yaml
from integration._support import prompt_cache_breakpoint as pcb
from integration._support import responses_vendor as rv
from integration._support.client import Gateway, eventually, gateway_from_environment, string_value
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
_GLOBAL_UNSET: Final = "bridge-breakpoint-global-unset"
_GLOBAL_FALSE: Final = "bridge-breakpoint-global-false"
_ENDPOINTS: Final = ("chat", "messages", "responses")
Endpoint: TypeAlias = Literal["chat", "messages", "responses"]
_RecordProperty: TypeAlias = Callable[[str, object], None]
@dataclass(frozen=True, slots=True)
class _Call:
endpoint: Endpoint
stream: bool
marker: str
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
text: str
call_id: str
@dataclass(frozen=True, slots=True)
class _GlobalRig:
wire: Wire
proxy: OwnedProxy
@property
def gateway(self) -> Gateway:
return self.proxy.gateway
def _global_config(directory: Path, api_base: str) -> Path:
stock: Final = rv.JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
deployment: Final[Mapping[str, JsonValue]] = {
"model": pcb.MODEL,
"api_base": api_base,
"api_key": "integration-provider-key",
}
config: Final[Mapping[str, JsonValue]] = {
**stock,
"model_list": [
{"model_name": _GLOBAL_UNSET, "litellm_params": dict(deployment)},
{"model_name": _GLOBAL_FALSE, "litellm_params": {**deployment, "drop_params": False}},
],
"litellm_settings": {**rv.JSON_OBJECT.validate_python(stock["litellm_settings"]), "drop_params": True},
"router_settings": {**rv.JSON_OBJECT.validate_python(stock.get("router_settings") or {}), "num_retries": 0},
}
path: Final = directory / "bridge-breakpoint-global.yaml"
path.write_text(yaml.safe_dump(config))
return path
@pytest.fixture(scope="module")
def global_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_GlobalRig]:
directory: Final = tmp_path_factory.mktemp("bridge-breakpoint-global")
with wire_server(pcb.respond) as wire, gateway_from_environment() as gateway:
config: Final = _global_config(directory, f"{wire.url}/v1")
with owned_proxy_process(gateway, directory, {}, config=config, workers=2) as owned:
yield _GlobalRig(wire, owned)
@pytest.fixture(scope="module")
def spend() -> Iterator[pcb.SpendLogs]:
with pcb.spend_logs() as logs:
yield logs
def _path(endpoint: Endpoint) -> str:
match endpoint:
case "chat":
return "/v1/chat/completions"
case "messages":
return "/v1/messages"
case "responses":
return "/v1/responses"
def _body(model: str, call: _Call, breakpoint: JsonValue) -> Mapping[str, JsonValue]:
common: Final[Mapping[str, JsonValue]] = {"model": model, "stream": call.stream, **pcb.NO_CACHE}
text: Final = pcb.marked(pcb.text(pcb.prompt(call.marker)), breakpoint)
match call.endpoint:
case "chat":
return {**common, "messages": [{"role": "user", "content": [text]}]}
case "messages":
return {**common, "max_tokens": 64, "messages": [{"role": "user", "content": [text]}]}
case "responses":
return {
**common,
"input": [{"type": "message", "role": "user", "content": [{**text, "type": "input_text"}]}],
}
def _calls(count: int, endpoints: tuple[Endpoint, ...], *, stream: bool | None = None) -> tuple[_Call, ...]:
return tuple(
_Call(endpoints[index % len(endpoints)], index % 2 == 1 if stream is None else stream, uuid.uuid4().hex)
for index in range(count)
)
async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call, breakpoint: JsonValue) -> _Served:
async with client.stream(
"POST",
_path(call.endpoint),
json=_body(model, call, breakpoint),
headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"},
) as response:
raw: Final = await response.aread()
return _Served(call, response.status_code, raw.decode(), response.headers["x-litellm-call-id"])
async def _burst(
gateway: Gateway,
model: str,
calls: tuple[_Call, ...],
*,
breakpoint: JsonValue = pcb.EXPLICIT,
tolerate_transport_errors: bool = False,
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
results: Final = await asyncio.gather(
*(_send(client, gateway.key, model, call, breakpoint) for call in calls),
return_exceptions=tolerate_transport_errors,
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Served))
def _frames(text: str) -> tuple[Mapping[str, JsonValue], ...]:
return tuple(rv.JSON_OBJECT.validate_json(line[6:]) for line in text.splitlines() if line.startswith("data: {"))
def _upstream_id_shown_to_caller(served: _Served) -> str | None:
if served.call.endpoint == "messages":
return None
if not served.call.stream:
return string_value(rv.JSON_OBJECT.validate_json(served.text)["id"])
frames: Final = _frames(served.text)
if served.call.endpoint == "responses":
(completed,) = [frame for frame in frames if frame.get("type") == "response.completed"]
return string_value(rv.JSON_OBJECT.validate_python(completed["response"])["id"])
return string_value(frames[0]["id"])
def _assert_answered_in_its_own_shape(served: _Served) -> None:
assert served.status == 200, served.text
assert set(rv.MARKER.findall(served.text)) == {served.call.marker}, served.text
assert served.text.startswith(("event:", "data:")) == served.call.stream, served.text
assert served.text.startswith("{") != served.call.stream, served.text
assert ("response.completed" in served.text) == (served.call.stream and served.call.endpoint == "responses")
shown: Final = _upstream_id_shown_to_caller(served)
assert shown is None or pcb.answers(shown, served.call.marker), served.text
def _marked_once(posts: Sequence[Request], calls: Sequence[_Call], expected: JsonValue) -> None:
by_marker: Final = {marker: request for request in posts if (marker := rv.newest_marker(request.body.decode()))}
assert len(by_marker) == len(posts), [request.body for request in posts]
assert set(by_marker) == {call.marker for call in calls}, sorted(by_marker)
for call in calls:
block: Final = pcb.single_block(pcb.input_items(by_marker[call.marker]), "user")
assert block["type"] == "input_text" and block["text"] == pcb.prompt(call.marker), block
pcb.assert_marker(block, expected)
def _assert_each_lands_once(
spend: pcb.SpendLogs, model: str, failed: Sequence[_Served], served: Sequence[_Served]
) -> None:
expected: Final = len(failed) + len(served)
rows: Final = eventually(lambda: spend.rows_for(model), lambda found: len(found) >= expected, seconds=70)
by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows}
assert len(by_call) == len(rows) == expected, rows
for item in failed:
assert by_call[item.call_id]["status"] == "failure", (item.call_id, rows)
for item in served:
row: Final = by_call[item.call_id]
assert row["status"] == "success", (item.call_id, row)
shown: Final = _upstream_id_shown_to_caller(item)
assert shown is None or rv.same_response(string_value(row["request_id"]), shown), (row, shown)
def _health(gateway: Gateway, model: str) -> Mapping[str, JsonValue]:
response: Final = gateway.request("GET", f"/health?model={model}", None)
assert response.status_code in (200, 503), response.text
return rv.JSON_OBJECT.validate_json(response.text)
def _free_port() -> int:
with wire_server(pcb.respond) as probe:
port: Final = urlsplit(probe.url).port
assert port is not None, probe.url
return port
def _chat(gateway: Gateway, model: str, marker: str, breakpoint: JsonValue) -> httpx.Response:
return gateway.request("POST", "/v1/chat/completions", dict(_body(model, _Call("chat", False, marker), breakpoint)))
def _completion(response: httpx.Response, marker: str) -> str:
assert response.status_code == 200, response.text
body: Final = rv.JSON_OBJECT.validate_json(response.text)
assert pcb.answers(string_value(body["id"]), marker), body
assert rv.answer(marker) in response.text, response.text
return response.headers["x-litellm-call-id"]
def _user_block_on_wire(wire: Wire, marker: str) -> dict[str, JsonValue]:
block: Final = pcb.single_block(pcb.input_items(pcb.posted(wire, marker)), "user")
assert block["type"] == "input_text" and block["text"] == pcb.prompt(marker), block
return block
@pytest.mark.parametrize("model", (_GLOBAL_UNSET, _GLOBAL_FALSE), ids=("deployment-unset", "deployment-false"))
def test_global_drop_params_drops_a_malformed_marker(global_rig: _GlobalRig, model: str, spend: pcb.SpendLogs) -> None:
marker: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(global_rig.gateway, model, marker, "yes"), marker)
pcb.assert_marker(_user_block_on_wire(global_rig.wire, marker), None)
spend.landed(model, call_id, marker)
control: Final = uuid.uuid4().hex
control_id: Final = _completion(_chat(global_rig.gateway, model, control, pcb.EXPLICIT), control)
pcb.assert_marker(_user_block_on_wire(global_rig.wire, control), pcb.EXPLICIT)
spend.landed(model, control_id, control)
async def test_mixed_burst_carries_every_marker_once(gateway: Gateway, spend: pcb.SpendLogs) -> None:
calls: Final = _calls(24, _ENDPOINTS)
with wire_server(pcb.respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True)
served: Final = await _burst(gateway, model, calls)
assert len(served) == 24
for item in served:
_assert_answered_in_its_own_shape(item)
_marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT)
_assert_each_lands_once(spend, model, (), served)
async def test_upstream_outage_fails_cleanly_and_the_restarted_upstream_serves_marked_calls(
gateway: Gateway, spend: pcb.SpendLogs
) -> None:
port: Final = _free_port()
while_down: Final = _calls(12, _ENDPOINTS)
after: Final = _calls(12, _ENDPOINTS)
with gateway.scenario() as scenario:
model: Final = scenario.model(model=pcb.MODEL, api_base=f"http://127.0.0.1:{port}/v1", drop_params=True)
failed: Final = await _burst(gateway, model, while_down)
assert len(failed) == 12
for item in failed:
assert item.status >= 500, (item.status, item.text)
assert "answer marker" not in item.text and "event:" not in item.text, item.text
down: Final = _health(gateway, model)
assert (down["healthy_count"], down["unhealthy_count"]) == (0, 1), down
with wire_server(pcb.respond, port=port) as wire:
_health(gateway, model)
probes: Final = pcb.drained_posts(wire)
assert [rv.newest_marker(request.body.decode()) for request in probes] == [None], probes
served: Final = await _burst(gateway, model, after)
assert len(served) == 12
for item in served:
_assert_answered_in_its_own_shape(item)
_marked_once(pcb.drained_posts(wire), after, pcb.EXPLICIT)
_assert_each_lands_once(spend, model, failed, served)
def _slow(request: Request) -> Reply:
reply: Final = pcb.respond(request)
return dataclasses.replace(reply, pause_between_chunks=0.4) if reply.chunks else reply
async def test_concurrent_slow_streams_each_complete_with_one_upstream_call(
gateway: Gateway, spend: pcb.SpendLogs
) -> None:
calls: Final = _calls(6, ("chat",), stream=True)
with wire_server(_slow) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True)
served: Final = await _burst(gateway, model, calls)
assert len(served) == 6
for item in served:
_assert_answered_in_its_own_shape(item)
_marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT)
_assert_each_lands_once(spend, model, (), served)
@dataclass(frozen=True, slots=True)
class _Held:
release: threading.Event
markers: SimpleQueue[str]
def respond(self, request: Request) -> Reply:
marker: Final = rv.newest_marker(request.body.decode()) if request.method == "POST" else None
if marker is None:
return pcb.respond(request)
self.markers.put(marker)
if not self.release.wait(timeout=60):
return rv.error(504, "the burst was never released", "held")
return pcb.respond(request)
def _worker_pids(owned: OwnedProxy) -> tuple[int, ...]:
return eventually(lambda: pcb.started_worker_pids(owned.log), lambda pids: len(pids) == 2, seconds=30)
async def _hold_burst(
held: _Held, candidate: Gateway, model: str, calls: tuple[_Call, ...]
) -> asyncio.Task[tuple[_Served, ...]]:
burst: Final = asyncio.create_task(_burst(candidate, model, calls, tolerate_transport_errors=True))
await asyncio.to_thread(eventually, held.markers.qsize, lambda size: size == len(calls), 60)
return burst
async def test_worker_sigkill_mid_burst_leaves_the_sibling_answering(
gateway: Gateway, tmp_path: Path, spend: pcb.SpendLogs
) -> None:
calls: Final = _calls(20, ("chat",), stream=False)
held: Final = _Held(threading.Event(), SimpleQueue())
with wire_server(held.respond) as wire:
config: Final = _global_config(tmp_path, f"{wire.url}/v1")
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
workers: Final = _worker_pids(owned)
burst: Final = await _hold_burst(held, candidate, _GLOBAL_UNSET, calls)
held_by: Final = MappingProxyType({pid: pcb.open_upstream_connections(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == 20, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
held.release.set()
served: Final = await burst
assert held_by[survivor_pid] >= 10, held_by
assert len(served) == held_by[survivor_pid], (held_by, len(served))
for item in served:
_assert_answered_in_its_own_shape(item)
_marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT)
follow_up: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(candidate, _GLOBAL_UNSET, follow_up, "yes"), follow_up)
pcb.assert_marker(_user_block_on_wire(wire, follow_up), None)
spend.landed(_GLOBAL_UNSET, call_id, follow_up)
async def test_proxy_restart_mid_burst_never_lands_a_served_call_twice(
gateway: Gateway, tmp_path: Path, record_property: _RecordProperty, spend: pcb.SpendLogs
) -> None:
calls: Final = _calls(20, ("chat",), stream=False)
held: Final = _Held(threading.Event(), SimpleQueue())
with wire_server(held.respond) as wire:
config: Final = _global_config(tmp_path, f"{wire.url}/v1")
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first:
_worker_pids(first)
burst: Final = await _hold_burst(held, first.gateway, _GLOBAL_UNSET, calls)
first.process.terminate()
held.release.set()
served: Final = await burst
for item in served:
_assert_answered_in_its_own_shape(item)
second_directory: Final = tmp_path / "second"
second_directory.mkdir()
with owned_proxy_process(gateway, second_directory, {}, config=config, workers=2) as second:
follow_up: Final = uuid.uuid4().hex
call_id: Final = _completion(_chat(second.gateway, _GLOBAL_UNSET, follow_up, pcb.EXPLICIT), follow_up)
pcb.assert_marker(_user_block_on_wire(wire, follow_up), pcb.EXPLICIT)
spend.landed(_GLOBAL_UNSET, call_id, follow_up)
counts: Final = Counter(string_value(row["litellm_call_id"]) for row in spend.rows_for(_GLOBAL_UNSET))
assert all(count == 1 for count in counts.values()), counts
landed: Final = sum(1 for item in served if item.call_id in counts)
record_property("served", len(served))
record_property("landed", landed)
record_property("lost_responses", len(calls) - len(served))

View file

@ -136,9 +136,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag
function_call_output = item
break
assert (
function_call_output is not None
), "function_call_output not found in response"
assert function_call_output is not None, "function_call_output not found in response"
assert function_call_output["call_id"] == "call_abc123"
# Check that the output is correctly transformed
@ -148,12 +146,8 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag
image_item = output[0]
# Should be transformed to Responses API format
assert (
image_item["type"] == "input_image"
), f"Expected type 'input_image', got '{image_item.get('type')}'"
assert (
image_item["image_url"] == test_image_base64
), "image_url should be a flat string, not a nested object"
assert image_item["type"] == "input_image", f"Expected type 'input_image', got '{image_item.get('type')}'"
assert image_item["image_url"] == test_image_base64, "image_url should be a flat string, not a nested object"
assert "detail" in image_item, "detail field should be present"
print("✓ Tool result with image correctly transformed to Responses API format")
@ -215,9 +209,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text
function_call_output = item
break
assert (
function_call_output is not None
), "function_call_output not found in response"
assert function_call_output is not None, "function_call_output not found in response"
assert function_call_output["call_id"] == "call_abc123"
# Check that the output is correctly transformed to use input_text, not output_text
@ -227,16 +219,12 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text
text_item = output[0]
# Should be transformed to use input_text for tool results in Responses API format
assert (
text_item["type"] == "input_text"
), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'"
assert (
text_item["text"] == "15 degrees"
), f"Expected text '15 degrees', got '{text_item.get('text')}'"
print(
"✓ Tool result with text correctly transformed to use input_text for Responses API format"
assert text_item["type"] == "input_text", (
f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'"
)
assert text_item["text"] == "15 degrees", f"Expected text '15 degrees', got '{text_item.get('text')}'"
print("✓ Tool result with text correctly transformed to use input_text for Responses API format")
def test_openai_responses_chunk_parser_reasoning_summary():
@ -245,9 +233,7 @@ def test_openai_responses_chunk_parser_reasoning_summary():
)
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"delta": "**Compar",
@ -279,9 +265,7 @@ def test_chunk_parser_string_output_text_delta_produces_text():
)
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {"type": "response.output_text.delta", "delta": "literal text"}
@ -302,9 +286,7 @@ def test_chunk_parser_enum_output_text_delta_produces_text():
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {"type": ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, "delta": "enum text"}
@ -325,9 +307,7 @@ def test_chunk_parser_function_call_added_produces_tool_use():
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
@ -412,9 +392,7 @@ Tomorrow will bring its petitions and promises,
but for now the city breathes slow and wide,
and I learn to carry this small calm home."""
output_text = ResponseOutputText(
annotations=[], text=poem_text, type="output_text", logprobs=[]
)
output_text = ResponseOutputText(annotations=[], text=poem_text, type="output_text", logprobs=[])
output_message = ResponseOutputMessage(
id="msg_04c8021b8b3188a00068e9ae0b92f4819dac64d85b4abb67ec",
content=[output_text],
@ -426,9 +404,7 @@ and I learn to carry this small calm home."""
# Create usage information
usage = ResponseAPIUsage(
input_tokens=16,
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
output_tokens=195,
output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
total_tokens=211,
@ -777,11 +753,7 @@ def test_recover_output_items_merges_text_only_items_at_distinct_indices():
]
)
recovered = (
LiteLLMResponsesTransformationHandler._recover_output_items_from_raw_sse(
raw_sse
)
)
recovered = LiteLLMResponsesTransformationHandler._recover_output_items_from_raw_sse(raw_sse)
assert len(recovered) == 2
assert recovered[0]["id"] == "msg_item_0"
@ -919,9 +891,7 @@ def test_transform_request_system_only_message_maps_to_system_input_item():
{
"type": "message",
"role": "system",
"content": [
{"type": "input_text", "text": "You are a helpful assistant."}
],
"content": [{"type": "input_text", "text": "You are a helpful assistant."}],
}
]
# System content lives in input only; not duplicated into instructions.
@ -993,9 +963,7 @@ def test_transform_request_single_char_keys_not_matched():
assert result_correct.get("metadata") == {"user_id": "123"}
assert result_correct.get("previous_response_id") == "resp_abc"
print(
"✓ Single-character keys are not incorrectly matched to metadata/previous_response_id"
)
print("✓ Single-character keys are not incorrectly matched to metadata/previous_response_id")
# =============================================================================
@ -1015,9 +983,7 @@ def test_message_done_does_not_emit_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.output_item.done",
@ -1029,9 +995,9 @@ def test_message_done_does_not_emit_is_finished():
# After the fix, message completion should NOT set finish_reason
# ModelResponseStream doesn't have is_finished - check finish_reason instead
assert len(result.choices) > 0, "result should have choices"
assert (
result.choices[0].finish_reason is None or result.choices[0].finish_reason == ""
), "message completion should not emit finish_reason"
assert result.choices[0].finish_reason is None or result.choices[0].finish_reason == "", (
"message completion should not emit finish_reason"
)
def test_response_completed_emits_is_finished():
@ -1043,9 +1009,7 @@ def test_response_completed_emits_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {"type": "response.completed"}
@ -1053,9 +1017,7 @@ def test_response_completed_emits_is_finished():
# response.completed should emit finish_reason='stop'
assert len(result.choices) > 0, "result should have choices"
assert (
result.choices[0].finish_reason == "stop"
), "response.completed should emit finish_reason='stop'"
assert result.choices[0].finish_reason == "stop", "response.completed should emit finish_reason='stop'"
def test_response_completed_with_function_calls_emits_tool_calls_finish_reason():
@ -1074,9 +1036,7 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason()
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
# Simulate a response.completed event with function_call in output
# This matches what Azure/OpenAI sends for gpt-5.1-codex-mini and similar models
@ -1102,9 +1062,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason()
# response.completed with function_call should emit finish_reason='tool_calls'
assert len(result.choices) > 0, "result should have choices"
assert (
result.choices[0].finish_reason == "tool_calls"
), "response.completed with function_call output should emit finish_reason='tool_calls'"
assert result.choices[0].finish_reason == "tool_calls", (
"response.completed with function_call output should emit finish_reason='tool_calls'"
)
def test_response_completed_with_message_only_emits_stop_finish_reason():
@ -1115,9 +1075,7 @@ def test_response_completed_with_message_only_emits_stop_finish_reason():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
# Simulate a response.completed event with only message output
chunk = {
@ -1141,9 +1099,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason():
# response.completed with only message should emit finish_reason='stop'
assert len(result.choices) > 0, "result should have choices"
assert (
result.choices[0].finish_reason == "stop"
), "response.completed with only message output should emit finish_reason='stop'"
assert result.choices[0].finish_reason == "stop", (
"response.completed with only message output should emit finish_reason='stop'"
)
def test_response_completed_preserves_usage_with_cached_tokens():
@ -1159,9 +1117,7 @@ def test_response_completed_preserves_usage_with_cached_tokens():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.completed",
@ -1190,18 +1146,12 @@ def test_response_completed_preserves_usage_with_cached_tokens():
result = iterator.chunk_parser(chunk)
assert result.usage is not None, "usage should be set on response.completed chunk"
assert (
result.usage.prompt_tokens == 1226
), "prompt_tokens should map from input_tokens"
assert (
result.usage.completion_tokens == 5
), "completion_tokens should map from output_tokens"
assert (
result.usage.prompt_tokens_details is not None
), "prompt_tokens_details should be set"
assert (
result.usage.prompt_tokens_details.cached_tokens == 1024
), "cached_tokens should be preserved from input_tokens_details"
assert result.usage.prompt_tokens == 1226, "prompt_tokens should map from input_tokens"
assert result.usage.completion_tokens == 5, "completion_tokens should map from output_tokens"
assert result.usage.prompt_tokens_details is not None, "prompt_tokens_details should be set"
assert result.usage.prompt_tokens_details.cached_tokens == 1024, (
"cached_tokens should be preserved from input_tokens_details"
)
def test_function_call_done_emits_is_finished():
@ -1215,9 +1165,7 @@ def test_function_call_done_emits_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.output_item.done",
@ -1237,9 +1185,9 @@ def test_function_call_done_emits_is_finished():
"output_item.done for function_call must not emit finish_reason; "
"response.completed is responsible for the terminal finish_reason"
)
assert not result.choices[
0
].delta.tool_calls, "output_item.done for function_call must not include a duplicate tool_calls delta"
assert not result.choices[0].delta.tool_calls, (
"output_item.done for function_call must not include a duplicate tool_calls delta"
)
def test_text_plus_tool_calls_sequence():
@ -1254,9 +1202,7 @@ def test_text_plus_tool_calls_sequence():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
# Simulate the sequence from OpenAI Responses API
chunks = [
@ -1295,28 +1241,23 @@ def test_text_plus_tool_calls_sequence():
# Check message done (index 2) does NOT have finish_reason set
message_done_result = results[2]
assert len(message_done_result.choices) > 0, "message done should have choices"
assert (
message_done_result.choices[0].finish_reason is None
or message_done_result.choices[0].finish_reason == ""
), "message done should not have finish_reason"
assert message_done_result.choices[0].finish_reason is None or message_done_result.choices[0].finish_reason == "", (
"message done should not have finish_reason"
)
# Check function_call done (index 5) does NOT have finish_reason set
# (response.completed is responsible for the terminal finish_reason)
function_done_result = results[5]
assert (
len(function_done_result.choices) > 0
), "function_call done should have choices"
assert (
function_done_result.choices[0].finish_reason is None
), "output_item.done for function_call must not emit finish_reason"
assert len(function_done_result.choices) > 0, "function_call done should have choices"
assert function_done_result.choices[0].finish_reason is None, (
"output_item.done for function_call must not emit finish_reason"
)
# Check response.completed (index 6) has finish_reason='stop'
# (the mock chunk has no nested 'response' data, so has_function_calls is False → 'stop')
completed_result = results[6]
assert len(completed_result.choices) > 0, "response.completed should have choices"
assert (
completed_result.choices[0].finish_reason == "stop"
), "response.completed should have finish_reason='stop'"
assert completed_result.choices[0].finish_reason == "stop", "response.completed should have finish_reason='stop'"
# =============================================================================
@ -1333,7 +1274,11 @@ def test_developer_message_content_uses_input_text():
assert instructions is None
assert input_items == [
{"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "Always answer in French."}]}
{
"type": "message",
"role": "developer",
"content": [{"type": "input_text", "text": "Always answer in French."}],
}
]
@ -1395,9 +1340,7 @@ def test_tool_message_output_uses_input_text_not_output_text():
output = function_call_output["output"]
assert isinstance(output, list), f"output should be a list, got {type(output)}"
assert len(output) == 1
assert (
output[0]["type"] == "input_text"
), f"Expected input_text, got {output[0].get('type')}"
assert output[0]["type"] == "input_text", f"Expected input_text, got {output[0].get('type')}"
assert output[0]["text"] == '{"temperature": 15, "condition": "sunny"}'
print("✓ Tool message output correctly uses input_text type")
@ -1582,13 +1525,9 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
assert result is not None, f"Result should not be None for effort={effort}"
assert result["effort"] == effort, f"Effort should be {effort}"
assert (
"summary" not in result
), f"Summary should NOT be present by default for effort={effort}"
assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}"
print(
f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)"
)
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)")
# Test 2: With flag enabled - summary IS added
litellm.reasoning_auto_summary = True
@ -1598,9 +1537,9 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
assert result is not None, f"Result should not be None for effort={effort}"
assert result["effort"] == effort, f"Effort should be {effort}"
assert (
result["summary"] == "detailed"
), f"Summary should be 'detailed' when flag is enabled for effort={effort}"
assert result["summary"] == "detailed", (
f"Summary should be 'detailed' when flag is enabled for effort={effort}"
)
print(
f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)"
@ -1611,9 +1550,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true")
result = handler.map_reasoning_effort("high")
assert (
result["summary"] == "detailed"
), "Summary should be 'detailed' when env var is enabled"
assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled"
print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly")
# Test 4: Dict input is passed through as-is (no modification)
@ -1627,9 +1564,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
assert result_dict["summary"] == "custom_summary"
print("✓ Dict input is passed through without modification")
print(
"✓ All reasoning_effort behaviors work correctly with flag/env var control"
)
print("✓ All reasoning_effort behaviors work correctly with flag/env var control")
finally:
# Restore original values
@ -1705,9 +1640,7 @@ def test_transform_response_preserves_annotations():
# Create usage information
usage = ResponseAPIUsage(
input_tokens=10,
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
output_tokens=20,
output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
total_tokens=30,
@ -1794,13 +1727,9 @@ def test_transform_response_preserves_annotations():
assert choice.message.content == "Here is some information with citations."
# Check that annotations are preserved
assert hasattr(
choice.message, "annotations"
), "Message should have annotations attribute"
assert hasattr(choice.message, "annotations"), "Message should have annotations attribute"
assert choice.message.annotations is not None, "Annotations should not be None"
assert (
len(choice.message.annotations) == 2
), f"Expected 2 annotations, got {len(choice.message.annotations)}"
assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}"
# Verify annotation content
annotation1 = choice.message.annotations[0]
@ -1822,9 +1751,7 @@ def test_transform_response_preserves_annotations():
assert result.usage.completion_tokens == 20
assert result.usage.total_tokens == 30
print(
"✓ Annotations from Responses API are correctly preserved in Chat Completions format"
)
print("✓ Annotations from Responses API are correctly preserved in Chat Completions format")
def test_apply_patch_tool_call_converted_to_chat_completion_tool_call():
@ -1989,9 +1916,7 @@ def test_multi_tool_call_stream_no_premature_finish():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunks = [
# 0: response created
@ -2067,12 +1992,10 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[done_idx]
assert r is not None, f"{label}: chunk_parser must return a result"
assert len(r.choices) > 0, f"{label}: result must have choices"
assert (
r.choices[0].finish_reason is None
), f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)"
assert not r.choices[
0
].delta.tool_calls, (
assert r.choices[0].finish_reason is None, (
f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)"
)
assert not r.choices[0].delta.tool_calls, (
f"{label}: output_item.done must not include a duplicate tool_calls delta"
)
@ -2084,12 +2007,8 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[added_idx]
if r is not None and r.choices and r.choices[0].delta.tool_calls:
tc = r.choices[0].delta.tool_calls[0]
assert (
tc.function.name == expected_name
), f"output_item.added for {expected_name}: tool_call name mismatch"
assert (
tc.id == expected_call_id
), f"output_item.added for {expected_name}: call_id mismatch"
assert tc.function.name == expected_name, f"output_item.added for {expected_name}: tool_call name mismatch"
assert tc.id == expected_call_id, f"output_item.added for {expected_name}: call_id mismatch"
# 3. argument delta events (indices 2 and 5) should carry arguments
for delta_idx, expected_args, label in [
@ -2099,17 +2018,15 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[delta_idx]
if r is not None and r.choices and r.choices[0].delta.tool_calls:
tc = r.choices[0].delta.tool_calls[0]
assert (
tc.function.arguments == expected_args
), f"{label}: argument delta mismatch"
assert tc.function.arguments == expected_args, f"{label}: argument delta mismatch"
# 4. Only response.completed (index 7) emits the terminal finish_reason
completed_result = results[7]
assert completed_result is not None, "response.completed must return a result"
assert len(completed_result.choices) > 0, "response.completed must have choices"
assert (
completed_result.choices[0].finish_reason == "tool_calls"
), "response.completed with function_call outputs must emit finish_reason='tool_calls'"
assert completed_result.choices[0].finish_reason == "tool_calls", (
"response.completed with function_call outputs must emit finish_reason='tool_calls'"
)
# 5. No chunk before the last one should have finish_reason set
for idx, r in enumerate(results[:-1]):
@ -2119,9 +2036,7 @@ def test_multi_tool_call_stream_no_premature_finish():
f"— only response.completed should terminate the stream"
)
print(
"✓ Multi-tool-call stream completes without premature finish_reason termination"
)
print("✓ Multi-tool-call stream completes without premature finish_reason termination")
# =============================================================================
@ -2202,16 +2117,13 @@ def test_streaming_parallel_tool_calls_have_distinct_indices():
]
for chunk in chunks:
result = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk
)
result = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
expected_index = chunk["output_index"]
for choice in result.choices:
if choice.delta.tool_calls:
for tc in choice.delta.tool_calls:
assert tc.index == expected_index, (
f"Event {chunk['type']}: expected tool_call.index={expected_index}, "
f"got {tc.index}"
f"Event {chunk['type']}: expected tool_call.index={expected_index}, got {tc.index}"
)
@ -2339,9 +2251,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
},
]
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
results = [iterator.chunk_parser(chunk) for chunk in chunks]
# 1. output_item.done events (indices 4 and 8) must NOT emit finish_reason
@ -2353,9 +2263,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
f"{label}: output_item.done must not emit finish_reason "
f"(would prematurely terminate stream before subsequent tool calls arrive)"
)
assert not r.choices[
0
].delta.tool_calls, (
assert not r.choices[0].delta.tool_calls, (
f"{label}: output_item.done must not emit a duplicate tool_calls delta"
)
@ -2389,19 +2297,15 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
for tc in tool_calls:
if tc.function and tc.function.arguments:
idx = tc.index
assembled_args[idx] = (
assembled_args.get(idx, "") + tc.function.arguments
)
assembled_args[idx] = assembled_args.get(idx, "") + tc.function.arguments
# delta 1 = '{"path":' + delta 2 = '"/etc/foo"}' → '{"path":"/etc/foo"}'
assert assembled_args.get(0) == '{"path":"/etc/foo"}', (
f"Assembled args for index 0 (read_file): "
f"expected '{{\"path\":\"/etc/foo\"}}', got '{assembled_args.get(0)}'"
f"Assembled args for index 0 (read_file): expected '{{\"path\":\"/etc/foo\"}}', got '{assembled_args.get(0)}'"
)
# delta 1 = '{"path":' + delta 2 = '"/tmp"}' → '{"path":"/tmp"}'
assert assembled_args.get(1) == '{"path":"/tmp"}', (
f"Assembled args for index 1 (list_dir): "
f"expected '{{\"path\":\"/tmp\"}}', got '{assembled_args.get(1)}'"
f"Assembled args for index 1 (list_dir): expected '{{\"path\":\"/tmp\"}}', got '{assembled_args.get(1)}'"
)
# 4. Stream terminates with exactly one finish event, at the final response.completed chunk
@ -2410,16 +2314,13 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
for i, r in enumerate(results)
if r is not None and r.choices and r.choices[0].finish_reason
]
assert (
len(finish_events) == 1
), f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}"
assert len(finish_events) == 1, f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}"
assert finish_events[0][0] == len(chunks) - 1, (
f"Finish event must be at the last chunk (index {len(chunks) - 1}), "
f"but was at index {finish_events[0][0]}"
f"Finish event must be at the last chunk (index {len(chunks) - 1}), but was at index {finish_events[0][0]}"
)
assert finish_events[0][1] == "tool_calls", (
f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'"
)
assert (
finish_events[0][1] == "tool_calls"
), f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'"
# 5. Parallel tool calls have distinct indices matching output_index (0 and 1)
# Collect indices from output_item.added chunks only (they carry the call id)
@ -2435,9 +2336,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
1,
}, f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}"
print(
"✓ Parallel tool calls with split argument deltas stream correctly end-to-end"
)
print("✓ Parallel tool calls with split argument deltas stream correctly end-to-end")
def test_map_optional_params_preserves_reasoning_summary():
@ -2461,9 +2360,7 @@ def test_map_optional_params_preserves_reasoning_summary():
}
responses_api_request = ResponsesAPIOptionalRequestParams()
handler._map_optional_params_to_responses_api_request(
optional_params, responses_api_request
)
handler._map_optional_params_to_responses_api_request(optional_params, responses_api_request)
# Verify reasoning_effort dict with summary was fully preserved
assert "reasoning" in responses_api_request
@ -2736,9 +2633,7 @@ def test_reasoning_items_non_streaming_round_trip():
)
usage = ResponseAPIUsage(
input_tokens=10,
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
output_tokens=20,
output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
total_tokens=30,
@ -2802,9 +2697,7 @@ def test_reasoning_items_non_streaming_round_trip():
assert len(result.choices) == 1
msg = result.choices[0].message
assert (
msg.reasoning_content == summary_text
), "reasoning_content should equal summary text"
assert msg.reasoning_content == summary_text, "reasoning_content should equal summary text"
assert msg.reasoning_items is not None, "reasoning_items should be set"
assert len(msg.reasoning_items) == 1
@ -2829,13 +2722,9 @@ def test_reasoning_items_non_streaming_round_trip():
# The reasoning input item must appear before the assistant message item
types = [item.get("type") for item in input_items]
assert (
"reasoning" in types
), "reasoning input item must be emitted for the assistant turn"
assert "reasoning" in types, "reasoning input item must be emitted for the assistant turn"
reasoning_input = next(
item for item in input_items if item.get("type") == "reasoning"
)
reasoning_input = next(item for item in input_items if item.get("type") == "reasoning")
assert reasoning_input["id"] == "rs_test001"
assert reasoning_input["encrypted_content"] == encrypted
assert reasoning_input["summary"][0]["text"] == summary_text
@ -2843,13 +2732,9 @@ def test_reasoning_items_non_streaming_round_trip():
# reasoning item must come before the assistant message item
reasoning_idx = types.index("reasoning")
assistant_msg_idx = next(
i
for i, item in enumerate(input_items)
if item.get("type") == "message" and item.get("role") == "assistant"
i for i, item in enumerate(input_items) if item.get("type") == "message" and item.get("role") == "assistant"
)
assert (
reasoning_idx < assistant_msg_idx
), "reasoning input item must precede the assistant message item"
assert reasoning_idx < assistant_msg_idx, "reasoning input item must precede the assistant message item"
def test_reasoning_items_streaming_emitted_on_response_completed():
@ -2862,9 +2747,7 @@ def test_reasoning_items_streaming_emitted_on_response_completed():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
encrypted = "gAAAAABpw5xyz987FAKE=="
summary_text = "**Reasoning summary**\n\nModel thought about this carefully."
@ -2908,16 +2791,14 @@ def test_reasoning_items_streaming_emitted_on_response_completed():
assert result.choices[0].finish_reason == "stop"
# reasoning_items must be on the delta
assert (
getattr(delta, "reasoning_items", None) is not None
), "reasoning_items must be present on the response.completed delta"
assert getattr(delta, "reasoning_items", None) is not None, (
"reasoning_items must be present on the response.completed delta"
)
assert len(delta.reasoning_items) == 1
ri = delta.reasoning_items[0]
assert ri["type"] == "reasoning"
assert ri["id"] == "rs_stream001"
assert (
ri["encrypted_content"] == encrypted
), "encrypted_content must be preserved in streaming"
assert ri["encrypted_content"] == encrypted, "encrypted_content must be preserved in streaming"
assert ri["summary"][0]["text"] == summary_text
@ -2944,9 +2825,7 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id():
"arguments": "",
},
}
out = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk
)
out = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
tool_calls = out.model_dump()["choices"][0]["delta"]["tool_calls"]
assert tool_calls, "expected a tool_call chunk in the streaming delta"
return tool_calls[0]["id"]
@ -2965,9 +2844,7 @@ def test_streaming_chunks_share_one_chat_completion_id():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
events = [
{"type": "response.created", "response": {"id": "resp_abc", "output": []}},
{"type": "response.output_text.delta", "delta": "Hel"},
@ -2983,12 +2860,10 @@ def test_streaming_chunks_share_one_chat_completion_id():
assert len(set(ids)) == 1, f"streamed chunks carried different ids: {ids}"
assert ids[0], "streamed chunks carried an empty id"
other_stream = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
other_stream = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
assert other_stream.chunk_parser(events[1]).id != ids[0], (
"a separate stream must get its own id, not a process-wide one"
)
assert (
other_stream.chunk_parser(events[1]).id != ids[0]
), "a separate stream must get its own id, not a process-wide one"
@pytest.mark.asyncio
@ -2999,9 +2874,7 @@ def test_streaming_chunks_share_one_chat_completion_id():
({"include_usage": True}, None),
],
)
async def test_acompletion_bridge_normalizes_stream_options_on_the_wire(
stream_options, expected_wire_stream_options
):
async def test_acompletion_bridge_normalizes_stream_options_on_the_wire(stream_options, expected_wire_stream_options):
"""include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict."""
from unittest.mock import AsyncMock
@ -3077,9 +2950,7 @@ def test_chunk_parser_custom_tool_call_stream_sequence():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
added = iterator.chunk_parser(
{
@ -3157,9 +3028,7 @@ def test_chunk_parser_remaps_tool_call_indices_sequentially():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
first = iterator.chunk_parser(
{
@ -3736,9 +3605,7 @@ def _make_incomplete_responses_api_response(
created_at=1760144904,
error=None,
incomplete_details=(
{"reason": incomplete_reason}
if incomplete_reason is not None or empty_incomplete_details
else None
{"reason": incomplete_reason} if incomplete_reason is not None or empty_incomplete_details else None
),
instructions=None,
metadata={},
@ -3758,13 +3625,9 @@ def _make_incomplete_responses_api_response(
truncation="disabled",
usage=ResponseAPIUsage(
input_tokens=37,
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
output_tokens=16,
output_tokens_details=OutputTokensDetails(
reasoning_tokens=16, text_tokens=None
),
output_tokens_details=OutputTokensDetails(reasoning_tokens=16, text_tokens=None),
total_tokens=53,
cost=None,
),
@ -3814,9 +3677,7 @@ def _call_transform_response(
def test_transform_response_incomplete_reasoning_only_returns_empty_length_choice():
handler = LiteLLMResponsesTransformationHandler()
raw_response = _make_incomplete_responses_api_response(
"max_output_tokens", [_make_reasoning_only_output_item()]
)
raw_response = _make_incomplete_responses_api_response("max_output_tokens", [_make_reasoning_only_output_item()])
result = _call_transform_response(handler, raw_response)
@ -3835,9 +3696,7 @@ def test_transform_response_incomplete_reasoning_only_returns_empty_length_choic
def test_transform_response_incomplete_content_filter_maps_finish_reason():
handler = LiteLLMResponsesTransformationHandler()
raw_response = _make_incomplete_responses_api_response(
"content_filter", [_make_reasoning_only_output_item()]
)
raw_response = _make_incomplete_responses_api_response("content_filter", [_make_reasoning_only_output_item()])
result = _call_transform_response(handler, raw_response)
@ -3860,11 +3719,7 @@ def test_transform_response_completed_with_reasonless_incomplete_details_keeps_s
handler = LiteLLMResponsesTransformationHandler()
output_message = ResponseOutputMessage(
id="msg_complete",
content=[
ResponseOutputText(
annotations=[], text="full answer", type="output_text", logprobs=[]
)
],
content=[ResponseOutputText(annotations=[], text="full answer", type="output_text", logprobs=[])],
role="assistant",
status="completed",
type="message",
@ -3886,11 +3741,7 @@ def test_transform_response_incomplete_partial_text_overrides_finish_reason_to_l
handler = LiteLLMResponsesTransformationHandler()
output_message = ResponseOutputMessage(
id="msg_partial",
content=[
ResponseOutputText(
annotations=[], text="partial answer", type="output_text", logprobs=[]
)
],
content=[ResponseOutputText(annotations=[], text="partial answer", type="output_text", logprobs=[])],
role="assistant",
status="incomplete",
type="message",
@ -3912,9 +3763,7 @@ def test_response_incomplete_stream_event_emits_length_and_usage():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.incomplete",
@ -3955,9 +3804,7 @@ def test_response_incomplete_stream_event_content_filter_maps_finish_reason():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.incomplete",
@ -3979,9 +3826,7 @@ def test_response_incomplete_stream_event_without_details_defaults_to_length():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
chunk = {
"type": "response.incomplete",
@ -4061,9 +3906,7 @@ def test_thinking_only_assistant_turn_still_sends_its_reasoning():
{
"role": "assistant",
"content": None,
"thinking_blocks": [
{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "sig1"}
],
"thinking_blocks": [{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "sig1"}],
},
{"role": "user", "content": "Why?"},
]
@ -4089,9 +3932,7 @@ def test_stored_reasoning_items_win_over_thinking_blocks():
"summary": [{"type": "summary_text", "text": "August in Denver is dry."}],
}
],
"thinking_blocks": [
{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "rs_real"}
],
"thinking_blocks": [{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "rs_real"}],
},
]
@ -4581,9 +4422,7 @@ def test_convert_chat_completion_messages_to_responses_api_drops_prompt_cache_br
{
"role": "tool",
"tool_call_id": "call_1",
"content": [
{"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint}
],
"content": [{"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint}],
},
],
)
@ -4897,3 +4736,104 @@ def test_every_bridged_chunk_after_response_created_carries_the_served_service_t
relayed = [iterator.chunk_parser(event).model_dump().get("service_tier") for event in events]
assert relayed == ["default"] * len(events), relayed
def test_convert_chat_completion_messages_to_responses_api_keeps_prompt_cache_breakpoint_on_unknown_block():
"""The hook marks the last block of its target message, so a message ending in a block the bridge
cannot map reaches the stringify path and has to keep the marker there."""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
breakpoint_marker = {"mode": "explicit"}
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "describe this"},
{
"type": "input_audio",
"input_audio": {"data": "Zm9v", "format": "wav"},
"prompt_cache_breakpoint": breakpoint_marker,
},
],
},
]
response, _ = handler.convert_chat_completion_messages_to_responses_api(
messages, keep_prompt_cache_breakpoints=True
)
content = response[0]["content"]
assert [block["type"] for block in content] == ["input_text", "input_text"]
assert content[1]["prompt_cache_breakpoint"] == breakpoint_marker
_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS: Final = (
{"type": "text", "text": "a string marker", "prompt_cache_breakpoint": "explicit"},
{"type": "text", "text": "an unknown mode", "prompt_cache_breakpoint": {"mode": "bogus"}},
{
"type": "input_audio",
"input_audio": {"data": "Zm9v", "format": "wav"},
"prompt_cache_breakpoint": ["explicit"],
},
{"type": "text", "text": "unsupported ttl", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "1h"}},
{"type": "text", "text": "unknown key", "prompt_cache_breakpoint": {"mode": "explicit", "scope": "all"}},
{"type": "text", "text": "supported ttl", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}},
{"type": "text", "text": "well formed", "prompt_cache_breakpoint": {"mode": "explicit"}},
)
def test_convert_chat_completion_messages_to_responses_api_drops_malformed_prompt_cache_breakpoint_under_drop_params():
"""OpenAI's Responses API answered "Supported values are: '30m'" for a 1h breakpoint ttl on 2026-10-07,
so an unsupported ttl drops the marker as a unit while an unknown key is dropped from a valid one."""
handler = LiteLLMResponsesTransformationHandler()
messages = [{"role": "user", "content": list(_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS)}]
response, _ = handler.convert_chat_completion_messages_to_responses_api(
messages, drop_params=True, keep_prompt_cache_breakpoints=True
)
content = response[0]["content"]
assert [block.get("prompt_cache_breakpoint") for block in content] == [
None,
None,
None,
None,
{"mode": "explicit"},
{"mode": "explicit", "ttl": "30m"},
{"mode": "explicit"},
]
assert all("prompt_cache_breakpoint" not in block for block in content[:4])
def test_convert_chat_completion_messages_to_responses_api_keeps_malformed_prompt_cache_breakpoint_by_default():
handler = LiteLLMResponsesTransformationHandler()
messages = [{"role": "user", "content": list(_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS)}]
response, _ = handler.convert_chat_completion_messages_to_responses_api(
messages, keep_prompt_cache_breakpoints=True
)
content = response[0]["content"]
assert [block["prompt_cache_breakpoint"] for block in content] == [
block["prompt_cache_breakpoint"] for block in _HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS
]
def test_transform_request_drop_params_in_litellm_params_gates_the_prompt_cache_breakpoint_carry():
handler = LiteLLMResponsesTransformationHandler()
messages = [{"role": "user", "content": [{"type": "text", "text": "hi", "prompt_cache_breakpoint": "explicit"}]}]
result = handler.transform_request(
model="gpt-6.1-sol",
messages=messages,
optional_params={},
litellm_params={"drop_params": True},
headers={},
litellm_logging_obj=Mock(),
)
assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0]