diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 9309df3300c..c8c6dfd106a 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -19,7 +19,7 @@ from openai.types.responses.response_input_param import ( from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam from openai.types.responses.tool_param import FunctionToolParam -from pydantic import BaseModel +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) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 78196e90003..f7e3d644f34 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -652,6 +652,7 @@ class ChatCompletionCachedContent(TypedDict): class PromptCacheBreakpoint(TypedDict): mode: ReadOnly[Literal["explicit"]] + ttl: NotRequired[ReadOnly[Literal["30m"]]] class PromptCacheOptions(TypedDict, total=False): diff --git a/tests/integration/_support/prompt_cache_breakpoint.py b/tests/integration/_support/prompt_cache_breakpoint.py new file mode 100644 index 00000000000..ce3b97b86be --- /dev/null +++ b/tests/integration/_support/prompt_cache_breakpoint.py @@ -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 + ) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py new file mode 100644 index 00000000000..75bd9da7795 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py @@ -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) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py new file mode 100644 index 00000000000..7742d4b9821 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py @@ -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)) diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index eedf766ae92..5b2187211df 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -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]