mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(responses): keep cache breakpoints on blocks the bridge stringifies (#40032)
* fix(responses): keep cache breakpoints on requests that reach the Responses API The chat-completions bridge rebuilt every content block without the prompt_cache_breakpoint marker the cache-control hook had just placed on it, so the request went out with prompt_cache_options set to explicit mode and nothing actually marked. Explicit mode caches only what is marked, so those deployments lost the implicit caching they were getting before the injection point was added The native surface had a second miss. The bridging check looked the provider config up with the still-prefixed model while the config keys on the bare one, so every bedrock_mantle model read as having no native Responses support, and role-targeted injection points were deferred to a chat-completions pass that never runs for them * fix(responses): keep prompt cache breakpoint on stringified content blocks * test: cover the responses/ routing prefix in the native cache point case * fix(responses): keep only the fallback-branch marker carry, main already routes the hook #42281 landed the text, image_url and file branch carry and the implicit default on main, and the Responses routing lookup this branch changed has no observable effect there, so the PR shrinks to the unknown-block fallback branch and its regression test * fix(responses): drop a malformed prompt_cache_breakpoint under drop_params in the chat-to-Responses bridge * fix(responses): accept the 30m prompt_cache_breakpoint ttl under drop_params OpenAI's Responses API takes a ttl of 30m on an explicit breakpoint marker, so the drop_params validation keeps it instead of stripping it as an unknown key. Trims the bridge test docstrings to the dated vendor citation. * test(integration): audit cells for prompt_cache_breakpoint through the chat-to-Responses bridge Adds the /audit cells for the bridge's prompt_cache_breakpoint carry: valid and malformed markers on every carrying block kind and on tool and assistant content, the drop_params on and off contract, the hook-injected system marker, the Anthropic SDK path through the Responses adapter, upstream errors, idempotent spend rows, and the chaos cells (mixed burst, upstream outage, slow streams, worker SIGKILL and a proxy restart mid burst), all against a scripted Responses endpoint with the request it received read back by response id * test(integration): give every scripted responses reply its own id The prompt-cache-breakpoint audit's scripted upstream minted the response id from the request's marker, so three identical marked requests shared one id. The spend-log writer skips rows whose request_id already landed, and the idempotent-logging cell saw one row for three calls on every leg. The upstream now mints a unique id per reply like a real provider, and the cells match a response to its request by the marker inside that id.
This commit is contained in:
parent
c52082e10c
commit
9d3bf29d6d
6 changed files with 1383 additions and 297 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -652,6 +652,7 @@ class ChatCompletionCachedContent(TypedDict):
|
|||
|
||||
class PromptCacheBreakpoint(TypedDict):
|
||||
mode: ReadOnly[Literal["explicit"]]
|
||||
ttl: NotRequired[ReadOnly[Literal["30m"]]]
|
||||
|
||||
|
||||
class PromptCacheOptions(TypedDict, total=False):
|
||||
|
|
|
|||
210
tests/integration/_support/prompt_cache_breakpoint.py
Normal file
210
tests/integration/_support/prompt_cache_breakpoint.py
Normal file
|
|
@ -0,0 +1,210 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal, TypeAlias, assert_never
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import psutil
|
||||
import psycopg
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.client import eventually, object_value, string_value
|
||||
from integration._support.database import ROWS
|
||||
from integration._support.openai_wire import answering_model_discovery, responses_reply
|
||||
from integration._support.wire import Reply, Request, Wire
|
||||
from psycopg.rows import DictRow, dict_row
|
||||
from pydantic import JsonValue
|
||||
|
||||
MODEL: Final = "openai/responses/gpt-6.1-sol"
|
||||
EXPLICIT: Final[Mapping[str, JsonValue]] = {"mode": "explicit"}
|
||||
EXPLICIT_30M: Final[Mapping[str, JsonValue]] = {"mode": "explicit", "ttl": "30m"}
|
||||
NO_CACHE: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}}
|
||||
INJECTION: Final[Mapping[str, JsonValue]] = {
|
||||
"cache_control_injection_points": [{"location": "message", "role": "system"}],
|
||||
"prompt_cache_options": {"mode": "explicit"},
|
||||
}
|
||||
_SCRIPTED_FAILURE: Final = re.compile(r"fail-(\d{3})")
|
||||
_MINTED_RESPONSE: Final = re.compile(r"^resp_([0-9a-f]{32})-[0-9a-f]{32}$")
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
|
||||
Kind: TypeAlias = Literal["text", "image_url", "file", "input_audio"]
|
||||
KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "input_audio")
|
||||
WIRE_TYPE: Final[Mapping[Kind, str]] = {
|
||||
"text": "input_text",
|
||||
"image_url": "input_image",
|
||||
"file": "input_file",
|
||||
"input_audio": "input_text",
|
||||
}
|
||||
|
||||
|
||||
def _scripted(request: Request) -> Reply:
|
||||
body: Final = rv.JSON_OBJECT.validate_json(request.body)
|
||||
text: Final = request.body.decode()
|
||||
marker: Final = rv.newest_marker(text)
|
||||
failure: Final = _SCRIPTED_FAILURE.search(text)
|
||||
if failure is not None:
|
||||
return rv.error(int(failure.group(1)), f"scripted {failure.group(1)} marker-{marker}", "scripted_failure")
|
||||
return responses_reply(
|
||||
f"resp_{marker or uuid.uuid4().hex}-{uuid.uuid4().hex}",
|
||||
string_value(body["model"]),
|
||||
rv.answer(marker),
|
||||
stream=body.get("stream") is True,
|
||||
)
|
||||
|
||||
|
||||
respond: Final = answering_model_discovery(_scripted)
|
||||
|
||||
|
||||
def response_marker(identity: str) -> str | None:
|
||||
minted: Final = tuple(
|
||||
found for candidate in rv.response_identities(identity) if (found := _MINTED_RESPONSE.match(candidate))
|
||||
)
|
||||
return minted[0].group(1) if minted else None
|
||||
|
||||
|
||||
def answers(identity: str, marker: str) -> bool:
|
||||
return response_marker(identity) == marker
|
||||
|
||||
|
||||
def prompt(marker: str) -> str:
|
||||
return f"Say marker-{marker}"
|
||||
|
||||
|
||||
def text(value: str) -> dict[str, JsonValue]:
|
||||
return {"type": "text", "text": value}
|
||||
|
||||
|
||||
def marked(block: Mapping[str, JsonValue], marker: JsonValue) -> dict[str, JsonValue]:
|
||||
return {**block, "prompt_cache_breakpoint": marker}
|
||||
|
||||
|
||||
def block(kind: Kind, value: str) -> dict[str, JsonValue]:
|
||||
match kind:
|
||||
case "text":
|
||||
return text(value)
|
||||
case "image_url":
|
||||
return {"type": "image_url", "image_url": {"url": "https://example.com/breakpoint.png"}}
|
||||
case "file":
|
||||
return {"type": "file", "file": {"file_id": "file-breakpoint"}}
|
||||
case "input_audio":
|
||||
return {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}
|
||||
case _:
|
||||
assert_never(kind)
|
||||
|
||||
|
||||
def drained_posts(wire: Wire) -> tuple[Request, ...]:
|
||||
return tuple(request for request in wire.drain() if request.method == "POST")
|
||||
|
||||
|
||||
def with_marker(posts: Sequence[Request], marker: str) -> tuple[Request, ...]:
|
||||
return tuple(request for request in posts if f"marker-{marker}" in request.body.decode())
|
||||
|
||||
|
||||
def posted(wire: Wire, marker: str) -> Request:
|
||||
matching: Final = with_marker(drained_posts(wire), marker)
|
||||
assert len(matching) == 1, [request.body for request in matching]
|
||||
(request,) = matching
|
||||
assert request.target == "/v1/responses", request.target
|
||||
return request
|
||||
|
||||
|
||||
def body_of(request: Request) -> dict[str, JsonValue]:
|
||||
return rv.JSON_OBJECT.validate_json(request.body)
|
||||
|
||||
|
||||
def input_items(request: Request) -> list[dict[str, JsonValue]]:
|
||||
return rv.ITEMS.validate_python(body_of(request)["input"])
|
||||
|
||||
|
||||
def content_of(items: Sequence[Mapping[str, JsonValue]], role: str) -> list[dict[str, JsonValue]]:
|
||||
messages: Final = tuple(item for item in items if item.get("type") == "message" and item.get("role") == role)
|
||||
assert len(messages) == 1, items
|
||||
return rv.ITEMS.validate_python(messages[0]["content"])
|
||||
|
||||
|
||||
def single_block(items: Sequence[Mapping[str, JsonValue]], role: str) -> dict[str, JsonValue]:
|
||||
blocks: Final = content_of(items, role)
|
||||
assert len(blocks) == 1, blocks
|
||||
return blocks[0]
|
||||
|
||||
|
||||
def instruction_block(items: Sequence[Mapping[str, JsonValue]]) -> dict[str, JsonValue]:
|
||||
messages: Final = tuple(
|
||||
item for item in items if item.get("type") == "message" and item.get("role") in ("system", "developer")
|
||||
)
|
||||
assert len(messages) == 1, items
|
||||
blocks: Final = rv.ITEMS.validate_python(messages[0]["content"])
|
||||
assert len(blocks) == 1, blocks
|
||||
return blocks[0]
|
||||
|
||||
|
||||
def function_output(items: Sequence[Mapping[str, JsonValue]], call_id: str) -> list[dict[str, JsonValue]]:
|
||||
outputs: Final = tuple(
|
||||
item for item in items if item.get("type") == "function_call_output" and item.get("call_id") == call_id
|
||||
)
|
||||
assert len(outputs) == 1, items
|
||||
return rv.ITEMS.validate_python(outputs[0]["output"])
|
||||
|
||||
|
||||
def assert_marker(block_on_wire: Mapping[str, JsonValue], expected: JsonValue) -> None:
|
||||
if expected is None:
|
||||
assert "prompt_cache_breakpoint" not in block_on_wire, block_on_wire
|
||||
return
|
||||
assert block_on_wire.get("prompt_cache_breakpoint") == expected, block_on_wire
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpendLogs:
|
||||
connection: psycopg.Connection[DictRow]
|
||||
|
||||
def rows_for(self, model: str) -> list[dict[str, JsonValue]]:
|
||||
cursor: Final = self.connection.execute(
|
||||
'SELECT litellm_call_id, request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group = %s', (model,)
|
||||
)
|
||||
return ROWS.validate_python(cursor.fetchall())
|
||||
|
||||
def landed(
|
||||
self, model: str, call_id: str, marker: str | None, *, status: str = "success", seconds: float = 70
|
||||
) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: self.rows_for(model),
|
||||
lambda found: any(row["litellm_call_id"] == call_id for row in found),
|
||||
seconds=seconds,
|
||||
)
|
||||
matching: Final = tuple(row for row in rows if row["litellm_call_id"] == call_id)
|
||||
assert len(matching) == 1, rows
|
||||
(row,) = matching
|
||||
assert row["status"] == status, row
|
||||
assert marker is None or answers(string_value(row["request_id"]), marker), (row, marker)
|
||||
return row
|
||||
|
||||
|
||||
@contextmanager
|
||||
def spend_logs() -> Iterator[SpendLogs]:
|
||||
with psycopg.connect(os.environ["DATABASE_URL"], row_factory=dict_row, autocommit=True) as connection:
|
||||
connection.execute("SET default_transaction_read_only = on")
|
||||
yield SpendLogs(connection)
|
||||
|
||||
|
||||
def model_id(entries: Sequence[JsonValue], model: str) -> str:
|
||||
matching: Final = tuple(entry for entry in entries if object_value(entry).get("model_name") == model)
|
||||
assert len(matching) == 1, entries
|
||||
return string_value(object_value(object_value(matching[0])["model_info"])["id"])
|
||||
|
||||
|
||||
def started_worker_pids(log: Path) -> tuple[int, ...]:
|
||||
return tuple(int(found.group(1)) for found in _STARTED_WORKER.finditer(log.read_text()))
|
||||
|
||||
|
||||
def open_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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))
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue