diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index c8c6dfd106a..41f6a211036 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -5,6 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from itertools import accumulate, chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, Union, cast @@ -114,6 +115,83 @@ def _strip_prompt_cache_breakpoints(input_items: list[object]) -> list[object]: return [_strip_prompt_cache_breakpoints_from_item(item) for item in input_items] +def _is_audio_input_part(part: object) -> bool: + if not isinstance(part, dict): + return False + content_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return content_part.get("type") == "input_audio" + + +def _breakpoint_of(part: object) -> object: + if not isinstance(part, dict): + return None + content_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return content_part.get("prompt_cache_breakpoint") + + +def _pending_audio_breakpoint(pending: object, part: object) -> object: + if not _is_audio_input_part(part): + return None + return pending if pending is not None else _breakpoint_of(part) + + +def _carried_breakpoints(content: Sequence[object]) -> tuple[object, ...]: + seeded_from_the_end: Final = chain((None,), reversed(content)) + pending_after_each: Final = tuple(accumulate(seeded_from_the_end, _pending_audio_breakpoint)) + return tuple(reversed(pending_after_each[:-1])) + + +def _with_carried_breakpoint(part: object, marker: object) -> object: + if marker is None or not isinstance(part, dict) or _breakpoint_of(part) is not None: + return part + kept_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return {**kept_part, "prompt_cache_breakpoint": marker} + + +def _without_audio_input_parts_in_content(value: object) -> object: + if not isinstance(value, list): + return value + content: Final = cast(list[object], value) # cast-ok: isinstance confirms a list of content blocks + if not any(_is_audio_input_part(part) for part in content): + return value + return [ + _with_carried_breakpoint(part, marker) + for part, marker in zip(content, _carried_breakpoints(content)) + if not _is_audio_input_part(part) + ] + + +def _without_audio_input_parts_in_item(value: object) -> object: + if not isinstance(value, dict): + return value + input_item: Final = cast(dict[str, object], value) # cast-ok: isinstance confirms a Responses input item mapping + return { + key: _without_audio_input_parts_in_content(item) if key in ("content", "output") else item + for key, item in input_item.items() + } + + +def _without_audio_input_parts(input_items: list[object]) -> list[object]: + return [_without_audio_input_parts_in_item(item) for item in input_items] + + +def _supports_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool: + custom_llm_provider: Final = litellm_params.get("custom_llm_provider") + provider: Final = custom_llm_provider if isinstance(custom_llm_provider, str) else None + base_model: Final = litellm_params.get("base_model") + return litellm.supports_audio_input(model=model, custom_llm_provider=provider) or ( + isinstance(base_model, str) + and bool(base_model) + and litellm.supports_audio_input(model=base_model, custom_llm_provider=provider) + ) + + +def _drops_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool: + if not (litellm_params.get("drop_params") or litellm.drop_params): + return False + return not _supports_audio_input(model, litellm_params) + + def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]: return MappingProxyType( { @@ -678,6 +756,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # of instructions, mirroring how non-string system content is already # handled in convert_chat_completion_messages_to_responses_api. is_system_only_request: Final = not converted_input_items and converted_instructions is not None + target_input_items: Final = ( + _without_audio_input_parts(converted_input_items) + if _drops_audio_input(model, litellm_params) + else converted_input_items + ) input_items: Final = ( [ { @@ -687,7 +770,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): } ] if is_system_only_request - else converted_input_items + else target_input_items ) instructions: Final = None if is_system_only_request else converted_instructions @@ -1194,6 +1277,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) + elif original_type == "input_audio": + converted = with_prompt_cache_breakpoint( + {"type": "input_audio", "input_audio": item.get("input_audio")}, + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), + ) + result.append(converted) + verbose_logger.debug("Chat provider: input_audio -> %s", converted) else: # Try to map other types to responses API format item_type = original_type or "input_text" diff --git a/tests/integration/_support/prompt_cache_breakpoint.py b/tests/integration/_support/prompt_cache_breakpoint.py index ce3b97b86be..bfc7f1efbce 100644 --- a/tests/integration/_support/prompt_cache_breakpoint.py +++ b/tests/integration/_support/prompt_cache_breakpoint.py @@ -32,13 +32,13 @@ _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") +Kind: TypeAlias = Literal["text", "image_url", "file", "video_url"] +KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "video_url") WIRE_TYPE: Final[Mapping[Kind, str]] = { "text": "input_text", "image_url": "input_image", "file": "input_file", - "input_audio": "input_text", + "video_url": "input_text", } @@ -91,12 +91,19 @@ def block(kind: Kind, value: str) -> dict[str, JsonValue]: 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 "video_url": + return {"type": "video_url", "video_url": {"url": "https://example.com/clip.mp4"}} case _: assert_never(kind) +AUDIO_PAYLOAD: Final[Mapping[str, JsonValue]] = {"data": "Zm9v", "format": "wav"} + + +def audio() -> dict[str, JsonValue]: + return {"type": "input_audio", "input_audio": dict(AUDIO_PAYLOAD)} + + def drained_posts(wire: Wire) -> tuple[Request, ...]: return tuple(request for request in wire.drain() if request.method == "POST") diff --git a/tests/integration/providers/test_responses_bridge_input_audio_wire.py b/tests/integration/providers/test_responses_bridge_input_audio_wire.py new file mode 100644 index 00000000000..389d31dae3e --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_input_audio_wire.py @@ -0,0 +1,536 @@ +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 Reply, Request, Wire, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(120) + +Mode: TypeAlias = Literal["on", "off"] + +_MODES: Final[tuple[Mode, ...]] = ("on", "off") +_CAPABLE_MODEL: Final = "openai/responses/gpt-audio-mini" +_CAPABLE_WIRE_MODEL: Final = "gpt-audio-mini" +_HOSTILE: Final[tuple[tuple[str, JsonValue], ...]] = ( + ("int", 7), + ("list", ["Zm9v"]), + ("empty-string", ""), + ("5kb-string", "x" * 5000), +) +_HOSTILE_IDS: Final = tuple(name for name, _ in _HOSTILE) +_HOSTILE_VALUES: Final = tuple(value for _, value in _HOSTILE) + + +@dataclass(frozen=True, slots=True) +class _Bridge: + gateway: Gateway + wire: Wire + on: str + off: str + null: str + capable_on: str + base_model_param_on: str + base_model_info_on: str + injecting_off: str + spend: pcb.SpendLogs + + def model(self, mode: Mode) -> str: + return self.on if mode == "on" else self.off + + +@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=None), + scenario.model(model=_CAPABLE_MODEL, api_base=api_base, drop_params=True), + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True, base_model=_CAPABLE_WIRE_MODEL), + scenario.model( + model=pcb.MODEL, + api_base=api_base, + drop_params=True, + model_info={"base_model": _CAPABLE_WIRE_MODEL}, + ), + 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 _text(marker: str) -> dict[str, JsonValue]: + return {"type": "input_text", "text": pcb.prompt(marker)} + + +def _audio_on_wire(payload: JsonValue) -> dict[str, JsonValue]: + return {"type": "input_audio", "input_audio": payload} + + +def _user(marker: str, *parts: JsonValue) -> dict[str, JsonValue]: + return {"role": "user", "content": [pcb.text(pcb.prompt(marker)), *parts]} + + +def _chat(bridge: _Bridge, model: str, messages: Sequence[JsonValue], *, stream: bool = False) -> httpx.Response: + return bridge.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": list(messages), "stream": stream, **pcb.NO_CACHE}, + ) + + +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_request(bridge: _Bridge, marker: str, *, model: str = "gpt-6.1-sol", stream: bool = False) -> Request: + request: Final = pcb.posted(bridge.wire, marker) + body: Final = pcb.body_of(request) + assert body["model"] == model, body + assert (body.get("stream") is True) is stream, body + return request + + +def _user_content_on_wire( + bridge: _Bridge, marker: str, *, model: str = "gpt-6.1-sol", stream: bool = False +) -> list[dict[str, JsonValue]]: + return pcb.content_of(pcb.input_items(_wire_request(bridge, marker, model=model, stream=stream)), "user") + + +def _expected(mode: Mode, marker: str, *forwarded: JsonValue) -> list[JsonValue]: + return [_text(marker)] if mode == "on" else [_text(marker), *forwarded] + + +def test_audio_part_is_dropped_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_audio_part_is_forwarded_as_input_audio_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.off, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_openai_sdk_stream_shapes_the_audio_part_by_mode(bridge: _Bridge, mode: Mode) -> 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.model(mode), + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": pcb.prompt(marker)}, + {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}, + ], + } + ], + 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 + content: Final = _user_content_on_wire(bridge, marker, stream=True) + assert content == _expected(mode, marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("mode", _MODES) +async def test_async_openai_sdk_shapes_the_audio_part_by_mode(bridge: _Bridge, mode: Mode) -> 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.model(mode), + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": pcb.prompt(marker)}, + {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}, + ], + } + ], + 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 + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], marker) + + +def test_audio_capable_model_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.capable_on, [_user(marker, pcb.audio())]), marker) + content: Final = _user_content_on_wire(bridge, marker, model=_CAPABLE_WIRE_MODEL) + assert content == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))], content + bridge.spend.landed(bridge.capable_on, call_id, marker) + + +def test_base_model_in_litellm_params_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.base_model_param_on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.base_model_param_on, call_id, marker) + + +def test_base_model_in_model_info_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.base_model_info_on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.base_model_info_on, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_tool_output_audio_part_follows_the_mode(bridge: _Bridge, mode: Mode) -> 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": "record", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [pcb.text("heard it"), pcb.audio()]}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + output: Final = pcb.function_output(pcb.input_items(_wire_request(bridge, marker)), "call_1") + heard: Final[dict[str, JsonValue]] = {"type": "input_text", "text": "heard it"} + expected: Final[list[JsonValue]] = [heard] if mode == "on" else [heard, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + assert output == expected, output + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_injected_system_marker_lands_on_a_trailing_audio_part_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.audio()]} + call_id: Final = _completion(_chat(bridge, bridge.injecting_off, [system, _user(marker)]), marker) + request: Final = _wire_request(bridge, marker) + assert pcb.body_of(request)["prompt_cache_options"] == {"mode": "explicit"}, request.body + system_content: Final = pcb.content_of(pcb.input_items(request), "system") + assert system_content == [ + {"type": "input_text", "text": "sys"}, + pcb.marked(_audio_on_wire(dict(pcb.AUDIO_PAYLOAD)), pcb.EXPLICIT), + ], system_content + bridge.spend.landed(bridge.injecting_off, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_assistant_audio_part_follows_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": pcb.prompt(marker)}, + {"role": "assistant", "content": [pcb.text("earlier answer"), pcb.audio()]}, + {"role": "user", "content": "and again"}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + earlier: Final = pcb.content_of(pcb.input_items(_wire_request(bridge, marker)), "assistant") + spoken: Final[dict[str, JsonValue]] = {"type": "output_text", "text": "earlier answer"} + expected: Final[list[JsonValue]] = [spoken] if mode == "on" else [spoken, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + assert earlier == expected, earlier + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_anthropic_sdk_request_carries_no_audio_part_into_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + client: Final = anthropic.Anthropic( + base_url=str(bridge.gateway.client.base_url), api_key=bridge.gateway.key, max_retries=0 + ) + with client: + raw: Final = client.messages.with_raw_response.create( + model=bridge.on, + max_tokens=64, + messages=[{"role": "user", "content": [pcb.text(pcb.prompt(marker)), pcb.audio()]}], + extra_body=dict(pcb.NO_CACHE), + ) + (content,) = raw.parse().content + assert content.type == "text" and content.text == rv.answer(marker), content + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], None) + + +def test_native_responses_audio_part_never_enters_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + audio: Final = _audio_on_wire(dict(pcb.AUDIO_PAYLOAD)) + response: Final = bridge.gateway.request( + "POST", + "/v1/responses", + { + "model": bridge.on, + "input": [{"type": "message", "role": "user", "content": [_text(marker), audio]}], + **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 + assert _user_content_on_wire(bridge, marker) == [_text(marker), audio] + bridge.spend.landed(bridge.on, response.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("value", _HOSTILE_VALUES, ids=_HOSTILE_IDS) +def test_hostile_audio_value_is_forwarded_verbatim_without_drop_params(bridge: _Bridge, value: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + hostile: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": value} + call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, hostile)]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(value)] + bridge.spend.landed(bridge.off, call_id, marker) + + +@pytest.mark.parametrize("value", _HOSTILE_VALUES, ids=_HOSTILE_IDS) +def test_hostile_audio_value_is_dropped_under_drop_params(bridge: _Bridge, value: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + hostile: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": value} + call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, hostile)]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_audio_part_without_a_payload_follows_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + bare: Final[dict[str, JsonValue]] = {"type": "input_audio"} + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, bare)]), marker) + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, _audio_on_wire(None)), content + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_two_identical_audio_parts_follow_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, pcb.audio(), pcb.audio())]), marker) + audio: Final = _audio_on_wire(dict(pcb.AUDIO_PAYLOAD)) + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, audio, audio), content + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_audio_only_message_is_forwarded_with_empty_content_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": [pcb.audio()]}, + {"role": "assistant", "content": "I could not hear that"}, + {"role": "user", "content": pcb.prompt(marker)}, + ] + call_id: Final = _completion(_chat(bridge, bridge.on, messages), marker) + items: Final = pcb.input_items(_wire_request(bridge, marker)) + users: Final = [item for item in items if item.get("role") == "user"] + assert [item["content"] for item in users] == [[], [_text(marker)]], items + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_leading_audio_marker_has_no_preceding_part_to_carry_to(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.marked(pcb.audio(), pcb.EXPLICIT), pcb.text(pcb.prompt(marker))] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_each_dropped_audio_marker_moves_to_its_own_preceding_text(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.text(pcb.prompt(marker)), + pcb.marked(pcb.audio(), pcb.EXPLICIT), + pcb.text("second clip follows"), + pcb.marked(pcb.audio(), pcb.EXPLICIT_30M), + ] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [ + pcb.marked(_text(marker), pcb.EXPLICIT), + {"type": "input_text", "text": "second clip follows", "prompt_cache_breakpoint": pcb.EXPLICIT_30M}, + ] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_text_marker_wins_over_the_dropped_audio_marker(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.marked(pcb.text(pcb.prompt(marker)), pcb.EXPLICIT_30M), + pcb.marked(pcb.audio(), pcb.EXPLICIT), + ] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [pcb.marked(_text(marker), pcb.EXPLICIT_30M)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_null_drop_params_forwards_the_audio_part(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.null, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.null, call_id, marker) + + +@dataclass(frozen=True, slots=True) +class _Call: + mode: Mode + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + call_id: str + + +def _calls(count: int) -> tuple[_Call, ...]: + return tuple(_Call(_MODES[index % 2], index % 4 >= 2, uuid.uuid4().hex) for index in range(count)) + + +async def _send(client: httpx.AsyncClient, bridge: _Bridge, model: str, call: _Call) -> _Served: + body: Final[Mapping[str, JsonValue]] = { + "model": model, + "messages": [_user(call.marker, pcb.audio())], + "stream": call.stream, + **pcb.NO_CACHE, + } + async with client.stream( + "POST", "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {bridge.gateway.key}"} + ) as response: + raw: Final = await response.aread() + return _Served(call, response.status_code, raw.decode(), response.headers["x-litellm-call-id"]) + + +async def _burst(bridge: _Bridge, calls: Sequence[_Call], model_for: Mapping[Mode, str]) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=str(bridge.gateway.client.base_url), timeout=60, trust_env=False) as client: + return tuple(await asyncio.gather(*(_send(client, bridge, model_for[call.mode], call) for call in calls))) + + +def _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("data:") == served.call.stream, served.text + assert rv.answer(served.call.marker) in served.text, served.text + + +def _attempts_by_marker(bridge: _Bridge, calls: Sequence[_Call]) -> Mapping[str, tuple[Request, ...]]: + posts: Final = pcb.drained_posts(bridge.wire) + marked: Final = tuple((rv.newest_marker(request.body.decode()), request) for request in posts) + attempts: Final = { + call.marker: tuple(request for marker, request in marked if marker == call.marker) for call in calls + } + assert sum(len(group) for group in attempts.values()) == len(posts), [request.body for request in posts] + assert all(attempts.values()), sorted(marker for marker, group in attempts.items() if not group) + return attempts + + +def _posts_by_marker(bridge: _Bridge, calls: Sequence[_Call]) -> Mapping[str, Request]: + attempts: Final = _attempts_by_marker(bridge, calls) + assert all(len(group) == 1 for group in attempts.values()), { + marker: len(group) for marker, group in attempts.items() + } + return {marker: group[0] for marker, group in attempts.items()} + + +def _landed_once(bridge: _Bridge, model: str, served: Sequence[_Served], *, status: str = "success") -> None: + rows: Final = eventually( + lambda: bridge.spend.rows_for(model), + lambda found: {string_value(row["litellm_call_id"]) for row in found} >= {item.call_id for item in served}, + seconds=70, + ) + by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows} + assert len(by_call) == len(rows), rows + for item in served: + assert by_call[item.call_id]["status"] == status, (item.call_id, by_call[item.call_id]) + + +async def test_mixed_audio_burst_shapes_every_upstream_request_by_its_mode(bridge: _Bridge) -> None: + calls: Final = _calls(24) + served: Final = await _burst(bridge, calls, {"on": bridge.on, "off": bridge.off}) + assert len(served) == 24 + for item in served: + _answered_in_its_own_shape(item) + by_marker: Final = _posts_by_marker(bridge, calls) + for call in calls: + body: Final = pcb.body_of(by_marker[call.marker]) + assert (body.get("stream") is True) is call.stream, body + content: Final = pcb.content_of(pcb.input_items(by_marker[call.marker]), "user") + assert content == _expected(call.mode, call.marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + _landed_once(bridge, bridge.on, tuple(item for item in served if item.call.mode == "on")) + _landed_once(bridge, bridge.off, tuple(item for item in served if item.call.mode == "off")) + + +@dataclass(frozen=True, slots=True) +class _Doomed: + markers: frozenset[str] + + def respond(self, request: Request) -> Reply: + marker: Final = rv.newest_marker(request.body.decode()) if request.method == "POST" else None + if marker in self.markers: + return Reply(drop_connection=True) + return pcb.respond(request) + + +async def test_dropped_upstream_connections_fail_only_their_own_calls(bridge: _Bridge) -> None: + calls: Final = tuple(_Call("on", False, uuid.uuid4().hex) for _ in range(16)) + doomed: Final = _Doomed(frozenset(call.marker for index, call in enumerate(calls) if index % 4 == 0)) + with wire_server(doomed.respond) as wire, bridge.gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True) + rig: Final = _Bridge( + bridge.gateway, + wire, + model, + model, + model, + model, + model, + model, + model, + bridge.spend, + ) + served: Final = await _burst(rig, calls, {"on": model, "off": model}) + failed: Final = tuple(item for item in served if item.call.marker in doomed.markers) + answered: Final = tuple(item for item in served if item.call.marker not in doomed.markers) + assert (len(failed), len(answered)) == (4, 12), [(item.call.marker, item.status) for item in served] + for item in failed: + assert item.status >= 500, (item.status, item.text) + assert "answer marker" not in item.text, item.text + for item in answered: + _answered_in_its_own_shape(item) + attempts: Final = _attempts_by_marker(rig, calls) + for item in answered: + assert len(attempts[item.call.marker]) == 1, attempts[item.call.marker] + for call in calls: + for request in attempts[call.marker]: + assert pcb.content_of(pcb.input_items(request), "user") == [_text(call.marker)], request.body + _landed_once(rig, model, answered) + _landed_once(rig, model, failed, status="failure") 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 index 75bd9da7795..9fad3a2b191 100644 --- a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py @@ -234,6 +234,28 @@ def test_malformed_marker_is_dropped_on_every_block_kind(bridge: _Bridge, kind: bridge.spend.landed(bridge.on, call_id, marker) +def test_input_audio_block_is_forwarded_with_its_marker_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + call_id: Final = _completion(_chat(bridge, bridge.off, [{"role": "user", "content": content}]), marker) + second: Final = _second_block_on_wire(bridge, marker) + assert second["type"] == "input_audio" and second["input_audio"] == pcb.AUDIO_PAYLOAD, second + pcb.assert_marker(second, pcb.EXPLICIT) + bridge.spend.landed(bridge.off, call_id, marker) + + +def test_input_audio_block_is_dropped_under_drop_params_and_its_marker_moves_to_the_text(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_block_on_wire(bridge, marker) == { + "type": "input_text", + "text": pcb.prompt(marker), + "prompt_cache_breakpoint": pcb.EXPLICIT, + } + 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 @@ -274,17 +296,14 @@ def test_assistant_list_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValu 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", "")]} + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.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) + (system_block,) = pcb.content_of(pcb.input_items(request), "system") + assert system_block == {"type": "input_text", "text": "sys", "prompt_cache_breakpoint": pcb.EXPLICIT}, system_block bridge.spend.landed(bridge.injecting_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 index 7742d4b9821..6f9a05d8dfa 100644 --- 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 @@ -252,6 +252,26 @@ def test_global_drop_params_drops_a_malformed_marker(global_rig: _GlobalRig, mod spend.landed(model, control_id, control) +@pytest.mark.parametrize("model", (_GLOBAL_UNSET, _GLOBAL_FALSE), ids=("deployment-unset", "deployment-false")) +def test_global_drop_params_drops_the_audio_part_and_carries_its_marker( + global_rig: _GlobalRig, model: str, spend: pcb.SpendLogs +) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + response: Final = global_rig.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": content}], **pcb.NO_CACHE}, + ) + call_id: Final = _completion(response, marker) + assert _user_block_on_wire(global_rig.wire, marker) == { + "type": "input_text", + "text": pcb.prompt(marker), + "prompt_cache_breakpoint": pcb.EXPLICIT, + } + spend.landed(model, call_id, marker) + + 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: 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 5b2187211df..8c6c4059a22 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 @@ -3,6 +3,7 @@ import datetime import json import os import unittest +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast, get_args from unittest.mock import ANY, MagicMock, Mock, patch @@ -4738,6 +4739,188 @@ def test_every_bridged_chunk_after_response_created_carries_the_served_service_t assert relayed == ["default"] * len(events), relayed +_AUDIO_PART: Final = {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}} +_TEXT_PART: Final = {"type": "text", "text": "Transcribe this"} + + +def test_convert_chat_completion_messages_to_responses_api_maps_input_audio_block(): + handler: Final = LiteLLMResponsesTransformationHandler() + messages: Final = [ + { + "role": "user", + "content": [_TEXT_PART, {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}], + } + ] + + items, _ = handler.convert_chat_completion_messages_to_responses_api(messages, keep_prompt_cache_breakpoints=True) + + assert items[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + { + "type": "input_audio", + "input_audio": {"data": "Zm9v", "format": "wav"}, + "prompt_cache_breakpoint": {"mode": "explicit"}, + }, + ] + + +def test_convert_chat_completion_messages_to_responses_api_drops_malformed_input_audio_breakpoint_under_drop_params(): + handler: Final = LiteLLMResponsesTransformationHandler() + messages: Final = [{"role": "user", "content": [{**_AUDIO_PART, "prompt_cache_breakpoint": ["explicit"]}]}] + + items, _ = handler.convert_chat_completion_messages_to_responses_api( + messages, drop_params=True, keep_prompt_cache_breakpoints=True + ) + + assert items[0]["content"] == [_AUDIO_PART] + + +@pytest.fixture +def registered_audio_models(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setitem( + litellm.model_cost, + "unit-audio-capable", + {"litellm_provider": "openai", "mode": "chat", "supports_audio_input": True}, + ) + monkeypatch.setitem( + litellm.model_cost, + "unit-text-only", + { + "litellm_provider": "openai", + "mode": "chat", + "supports_audio_input": False, + "supports_prompt_cache_breakpoint": True, + }, + ) + + +def _bridge_input( + model: str, + drop_params: bool | None, + messages: Sequence[Mapping[str, object]] | None = None, + **extra_litellm_params: object, +) -> list[dict[str, object]]: + chat_messages: Final = cast( + List[AllMessageValues], list(messages or [{"role": "user", "content": [_TEXT_PART, _AUDIO_PART]}]) + ) # cast-ok: the tests build chat messages as plain mappings + request: Final = LiteLLMResponsesTransformationHandler().transform_request( + model=model, + messages=chat_messages, + optional_params={}, + litellm_params={"custom_llm_provider": "openai", "drop_params": drop_params, **extra_litellm_params}, + headers={}, + litellm_logging_obj=Mock(), + ) + return cast(list[dict[str, object]], request["input"]) # cast-ok: the bridge emits message item mappings + + +def test_transform_request_forwards_input_audio_without_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-text-only", drop_params=False)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_drops_input_audio_under_drop_params_when_model_lacks_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-text-only", drop_params=True)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"} + ] + + +def test_transform_request_keeps_input_audio_under_drop_params_when_model_supports_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-audio-capable", drop_params=True)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_keeps_input_audio_under_drop_params_when_base_model_supports_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("my-audio-deployment", drop_params=True, base_model="unit-audio-capable")[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_drops_input_audio_under_global_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", True) + + assert _bridge_input("unit-text-only", drop_params=None)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"} + ] + + +def test_transform_request_drops_input_audio_from_tool_output_under_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + {"role": "user", "content": "Describe the recording"}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "record", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [_TEXT_PART, _AUDIO_PART]}, + ] + + forwarded: Final = _bridge_input("unit-text-only", drop_params=False, messages=messages) + dropped: Final = _bridge_input("unit-text-only", drop_params=True, messages=messages) + + assert forwarded[-1]["type"] == "function_call_output" + assert forwarded[-1]["output"] == [{"type": "input_text", "text": "Transcribe this"}, _AUDIO_PART] + assert dropped[-1]["output"] == [{"type": "input_text", "text": "Transcribe this"}] + + +def test_transform_request_moves_the_dropped_audio_part_breakpoint_to_the_preceding_part( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + {"role": "user", "content": [_TEXT_PART, {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}]} + ] + + assert _bridge_input("unit-text-only", drop_params=True, messages=messages)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this", "prompt_cache_breakpoint": {"mode": "explicit"}} + ] + + +def test_transform_request_keeps_the_preceding_part_breakpoint_over_the_dropped_audio_part_breakpoint( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + { + "role": "user", + "content": [ + {**_TEXT_PART, "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}}, + {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}, + ], + } + ] + + assert _bridge_input("unit-text-only", drop_params=True, messages=messages)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}} + ] + + 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.""" @@ -4754,8 +4937,8 @@ def test_convert_chat_completion_messages_to_responses_api_keeps_prompt_cache_br "content": [ {"type": "text", "text": "describe this"}, { - "type": "input_audio", - "input_audio": {"data": "Zm9v", "format": "wav"}, + "type": "video_url", + "video_url": {"url": "https://example.com/clip.mp4"}, "prompt_cache_breakpoint": breakpoint_marker, }, ], @@ -4836,4 +5019,4 @@ def test_transform_request_drop_params_in_litellm_params_gates_the_prompt_cache_ litellm_logging_obj=Mock(), ) - assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0] + assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0] \ No newline at end of file