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

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 01:22:57 +00:00
parent ad6448d71b
commit 21a0b3c7d6
5 changed files with 797 additions and 734 deletions

View file

@ -617,7 +617,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
client: object | None = None,
) -> dict:
converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(messages)
supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model)
base_model: Final = litellm_params.get("base_model")
supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) or (
isinstance(base_model, str) and bool(base_model) and supports_openai_prompt_cache_breakpoint(base_model)
)
input_items_without_unsupported_markers: Final = (
converted_input_items
if supports_prompt_cache_breakpoint

View file

@ -0,0 +1,505 @@
from __future__ import annotations
import asyncio
import json
import re
import uuid
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from typing import Final, Literal, TypeAlias, cast
import httpx
from openai import AsyncOpenAI, OpenAI
from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolUnionParam
from pydantic import JsonValue, TypeAdapter
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request
from litellm.responses.utils import ResponsesAPIRequestUtils as _RU
_MODEL: Final = "openai/gpt-5.6"
_UNSUPPORTED_MODEL: Final = "openai/gpt-5.4-mini"
_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})")
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_BREAKPOINT: Final[dict[str, JsonValue]] = {"mode": "explicit"}
_TOOLS: Final[list[JsonValue]] = [
{
"type": "function",
"function": {
"name": "synthetic_tool",
"description": "Synthetic bridge test tool",
"parameters": {"type": "object", "properties": {}},
},
}
]
_IMAGE_URL: Final = "data:image/png;base64,aGVsbG8="
_ClientKind: TypeAlias = Literal["openai_sync", "openai_async", "httpx"]
_Surface: TypeAlias = Literal["chat", "responses"]
@dataclass(frozen=True, slots=True)
class _Call:
surface: _Surface
stream: bool
marker: str
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
response_id: str | None
text: str
def _response_id(marker: str) -> str:
return f"resp_{marker}"
def _request_marker(request: Request) -> str:
match: Final = _MARKER.search(request.body)
assert match is not None, request.body
return match.group(1).decode()
def _contains_breakpoint(value: JsonValue) -> bool:
if isinstance(value, dict):
return "prompt_cache_breakpoint" in value or any(_contains_breakpoint(item) for item in value.values())
if isinstance(value, list):
return any(_contains_breakpoint(item) for item in value)
return False
def _responses_body(marker: str) -> dict[str, JsonValue]:
response_id: Final = _response_id(marker)
return _JSON_OBJECT.validate_python(
{
"id": response_id,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"id": f"msg_{marker}",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": f"answer marker-{marker}", "annotations": []}],
}
],
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12},
}
)
def _responses_reply(request: Request, *, reject_breakpoints: bool = False) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=b'{"object":"list","data":[{"id":"gpt-5.6","object":"model"}]}')
body: Final = _JSON_OBJECT.validate_json(request.body)
if reject_breakpoints and _contains_breakpoint(body):
return Reply(
status=400,
body=json.dumps(
{
"error": {
"message": "prompt_cache_breakpoint is not supported on this model",
"type": "invalid_request_error",
"param": None,
"code": None,
}
}
).encode(),
)
marker: Final = _request_marker(request)
stream: Final = body.get("stream") is True
response: Final = _responses_body(marker)
if not stream:
return Reply(body=json.dumps(response).encode())
created: Final = {
"type": "response.created",
"sequence_number": 0,
"response": {**response, "status": "in_progress", "output": []},
}
delta: Final = {
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": f"msg_{marker}",
"output_index": 0,
"content_index": 0,
"delta": f"answer marker-{marker}",
}
completed: Final = {"type": "response.completed", "sequence_number": 2, "response": response}
events: Final = (created, delta, completed)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
)
def _prompt(marker: str, label: str) -> str:
return f"{label} marker-{marker}"
def _simple_chat_body(
model: str,
marker: str,
*,
stream: bool = False,
marked: bool = True,
system_as_string: bool = False,
prompt_cache_options: dict[str, JsonValue] | None = None,
) -> dict[str, JsonValue]:
marker_field: Final = {"prompt_cache_breakpoint": _BREAKPOINT} if marked else {}
user: Final = [{"type": "text", "text": _prompt(marker, "user"), **marker_field}]
messages: Final = (
[{"role": "system", "content": _prompt(marker, "system")}, {"role": "user", "content": user}]
if system_as_string
else [{"role": "user", "content": user}]
)
return _JSON_OBJECT.validate_python(
{
"model": model,
"messages": messages,
"tools": _TOOLS,
"reasoning_effort": "low",
"stream": stream,
"num_retries": 0,
**({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}),
}
)
def _multimodal_chat_body(model: str, marker: str, stream: bool) -> dict[str, JsonValue]:
return _JSON_OBJECT.validate_python(
{
"model": model,
"messages": [
{
"role": "system",
"content": [
{
"type": "text",
"text": _prompt(marker, "system"),
"prompt_cache_breakpoint": _BREAKPOINT,
}
],
},
{
"role": "user",
"content": [
{"type": "text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT},
{
"type": "image_url",
"image_url": {"url": _IMAGE_URL},
"prompt_cache_breakpoint": _BREAKPOINT,
},
{
"type": "file",
"file": {"file_id": "file-abc"},
"prompt_cache_breakpoint": _BREAKPOINT,
},
{"type": "text", "text": "unmarked extra text"},
],
},
],
"tools": _TOOLS,
"reasoning_effort": "low",
"stream": stream,
"num_retries": 0,
"prompt_cache_options": {"mode": "explicit"},
}
)
def _expected_multimodal_input(marker: str) -> list[JsonValue]:
return [
{
"type": "message",
"role": "system",
"content": [
{"type": "input_text", "text": _prompt(marker, "system"), "prompt_cache_breakpoint": _BREAKPOINT}
],
},
{
"type": "message",
"role": "user",
"content": [
{"type": "input_text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT},
{
"type": "input_image",
"image_url": _IMAGE_URL,
"detail": "auto",
"prompt_cache_breakpoint": _BREAKPOINT,
},
{"type": "input_file", "file_id": "file-abc", "prompt_cache_breakpoint": _BREAKPOINT},
{"type": "input_text", "text": "unmarked extra text"},
],
},
]
def _simple_expected_input(marker: str, *, marked: bool) -> list[JsonValue]:
text_block: Final = {"type": "input_text", "text": _prompt(marker, "user")}
return [
{
"type": "message",
"role": "user",
"content": [{**text_block, **({"prompt_cache_breakpoint": _BREAKPOINT} if marked else {})}],
}
]
def _request_body(request: Request) -> dict[str, JsonValue]:
assert request.method == "POST" and request.target == "/v1/responses", request.target
return _JSON_OBJECT.validate_json(request.body)
def _decoded_response_id(response_id: str) -> str:
decoded: Final = _RU._decode_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # reuse ID decoder
response_id
)
raw_response_id: Final = decoded.get("response_id")
assert isinstance(raw_response_id, str), decoded
return raw_response_id
def _spend_request_id_matches(
row: Mapping[str, JsonValue],
caller_response_id: str,
peer_response_id: str,
surface: _Surface,
) -> bool:
request_id: Final = row.get("request_id")
if not isinstance(request_id, str):
return False
match surface:
case "responses":
return request_id == caller_response_id
case "chat":
return _decoded_response_id(request_id) == peer_response_id
def _spend_rows(
model: str,
caller_response_id: str,
peer_response_id: str,
surface: _Surface,
) -> tuple[dict[str, JsonValue], ...]:
def matching_rows(rows: list[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]:
return tuple(
row for row in rows if _spend_request_id_matches(row, caller_response_id, peer_response_id, surface)
)
rows: Final = eventually(
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda candidates: len(matching_rows(candidates)) == 1,
seconds=60,
)
matched: Final = matching_rows(rows)
assert len(matched) == 1, matched
return matched
def _response_id_from_chat_stream(text: str) -> str:
payloads: Final = tuple(
_JSON_OBJECT.validate_json(line.removeprefix("data: "))
for line in text.splitlines()
if line.startswith("data: {")
)
assert payloads, text
response_id: Final = payloads[0].get("id")
assert isinstance(response_id, str), payloads[0]
return response_id
def _extra_body(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
return {
key: value
for key, value in body.items()
if key not in {"model", "messages", "tools", "reasoning_effort", "stream", "num_retries"}
}
def _sync_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served:
base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1"
model: Final = str(body["model"])
messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"])
tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"])
extras: Final = _extra_body(body)
with OpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client:
if stream:
response_stream: Final = client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=True,
extra_body=extras,
)
chunks: Final = tuple(response_stream)
assert chunks
return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "")
response: Final = client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=False,
extra_body=extras,
)
return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "")
async def _async_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served:
base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1"
model: Final = str(body["model"])
messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"])
tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"])
extras: Final = _extra_body(body)
async with AsyncOpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client:
if stream:
response_stream: Final = await client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=True,
extra_body=extras,
)
chunks: Final = tuple([chunk async for chunk in response_stream])
assert chunks
return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "")
response: Final = await client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=False,
extra_body=extras,
)
return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "")
async def _serve_chat(
gateway: Gateway,
body: dict[str, JsonValue],
client_kind: _ClientKind,
stream: bool,
) -> _Served:
match client_kind:
case "openai_sync":
return _sync_sdk_chat(gateway, body, stream)
case "openai_async":
return await _async_sdk_chat(gateway, body, stream)
case "httpx":
async with httpx.AsyncClient(
base_url=str(gateway.client.base_url),
headers={"Authorization": f"Bearer {gateway.key}"},
timeout=20,
trust_env=False,
) as client:
return await _raw_call(
client,
"/v1/chat/completions",
body,
_Call("chat", stream, _request_marker_from_body(body)),
)
def _request_marker_from_body(body: Mapping[str, JsonValue]) -> str:
match: Final = _MARKER.search(json.dumps(body).encode())
assert match is not None, body
return match.group(1).decode()
async def _raw_call(
client: httpx.AsyncClient,
path: str,
body: Mapping[str, JsonValue],
call: _Call,
) -> _Served:
async with client.stream(
"POST",
path,
json=body,
headers={"Authorization": f"Bearer {client.headers['Authorization'].removeprefix('Bearer ')}"},
) as response:
content: Final = await response.aread()
status: Final = response.status_code
text: Final = content.decode()
response_id: Final = (
_response_id_from_chat_stream(text)
if status == 200 and call.surface == "chat" and call.stream
else _JSON_OBJECT.validate_json(content).get("id")
if status == 200
else None
)
return _Served(call, status, response_id if isinstance(response_id, str) else None, text)
async def _send_call(
client: httpx.AsyncClient,
model: str,
call: _Call,
) -> _Served:
body: Final = (
_simple_chat_body(model, call.marker, stream=call.stream)
if call.surface == "chat"
else {
"model": model,
"input": _simple_expected_input(call.marker, marked=True),
"stream": call.stream,
"num_retries": 0,
}
)
path: Final = "/v1/chat/completions" if call.surface == "chat" else "/v1/responses"
try:
return await _raw_call(client, path, body, call)
except httpx.TransportError as error:
return _Served(call, 0, None, f"{type(error).__name__}: {error}")
async def _burst(
base_url: str,
key: str,
model: str,
calls: tuple[_Call, ...],
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(
base_url=base_url,
headers={"Authorization": f"Bearer {key}"},
timeout=20,
trust_env=False,
limits=httpx.Limits(max_connections=100),
) as client:
return tuple(await asyncio.gather(*(_send_call(client, model, call) for call in calls)))
def _calls(count: int) -> tuple[_Call, ...]:
surfaces: Final[tuple[_Surface, ...]] = ("chat", "chat", "responses")
return tuple(
_Call(
surface=surfaces[index % len(surfaces)],
stream=index % 3 == 1,
marker=uuid.uuid4().hex,
)
for index in range(count)
)
def _requests_for_marker(requests: tuple[Request, ...], marker: str) -> tuple[Request, ...]:
return tuple(
request
for request in requests
if request.method == "POST" and request.target == "/v1/responses" and _request_marker(request) == marker
)
def _peer_request_has_marker(request: Request, marker: str) -> bool:
body: Final = _request_body(request)
return _contains_breakpoint(body) and _request_marker(request) == marker
def _peer_marker_matches_response(served: _Served, requests: tuple[Request, ...]) -> bool:
peer_requests: Final = _requests_for_marker(requests, served.call.marker)
assert len(peer_requests) == 1, (served, peer_requests)
(peer_request,) = peer_requests
return _peer_request_has_marker(peer_request, served.call.marker)
def _assert_spend_for_result(served: _Served, model: str) -> None:
assert served.status == 200 and served.response_id is not None, served
peer_response_id: Final = _response_id(served.call.marker)
match served.call.surface:
case "responses":
(row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface)
case "chat":
assert _decoded_response_id(served.response_id) == peer_response_id, served
(row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface)
request_id: Final = row.get("request_id")
assert isinstance(request_id, str), row
match served.call.surface:
case "responses":
assert request_id == served.response_id, row
case "chat":
assert _decoded_response_id(request_id) == served.response_id, row

View file

@ -1,544 +1,33 @@
from __future__ import annotations
import asyncio
import json
import re
import signal
import threading
import uuid
from collections.abc import Iterable, Mapping
from contextlib import ExitStack
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from typing import Final, Literal, TypeAlias, cast
from urllib.parse import urlsplit
from typing import Final
import httpx
import psutil
import pytest
import yaml
from openai import AsyncOpenAI, OpenAI
from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolUnionParam
from pydantic import JsonValue, TypeAdapter
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from litellm.responses.utils import ResponsesAPIRequestUtils as _RU
_MODEL: Final = "openai/gpt-5.6"
_UNSUPPORTED_MODEL: Final = "openai/gpt-5.4-mini"
_CONFIG_MODEL: Final = "responses-bridge-cache-breakpoint-chaos"
_API_KEY: Final = "synthetic-responses-bridge-key"
_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})")
_STARTED_WORKER: Final[re.Pattern[str]] = re.compile(r"Started server process \[(\d+)\]")
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_BREAKPOINT: Final[dict[str, JsonValue]] = {"mode": "explicit"}
_TOOLS: Final[list[JsonValue]] = [
{
"type": "function",
"function": {
"name": "synthetic_tool",
"description": "Synthetic bridge test tool",
"parameters": {"type": "object", "properties": {}},
},
}
]
_IMAGE_URL: Final = "data:image/png;base64,aGVsbG8="
_ClientKind: TypeAlias = Literal["openai_sync", "openai_async", "httpx"]
_Surface: TypeAlias = Literal["chat", "responses"]
@dataclass(frozen=True, slots=True)
class _Call:
surface: _Surface
stream: bool
marker: str
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
response_id: str | None
text: str
def _response_id(marker: str) -> str:
return f"resp_{marker}"
def _request_marker(request: Request) -> str:
match: Final = _MARKER.search(request.body)
assert match is not None, request.body
return match.group(1).decode()
def _contains_breakpoint(value: JsonValue) -> bool:
if isinstance(value, dict):
return "prompt_cache_breakpoint" in value or any(_contains_breakpoint(item) for item in value.values())
if isinstance(value, list):
return any(_contains_breakpoint(item) for item in value)
return False
def _responses_body(marker: str) -> dict[str, JsonValue]:
response_id: Final = _response_id(marker)
return _JSON_OBJECT.validate_python(
{
"id": response_id,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"id": f"msg_{marker}",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": f"answer marker-{marker}", "annotations": []}],
}
],
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12},
}
)
def _responses_reply(request: Request, *, reject_breakpoints: bool = False) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=b'{"object":"list","data":[{"id":"gpt-5.6","object":"model"}]}')
body: Final = _JSON_OBJECT.validate_json(request.body)
if reject_breakpoints and _contains_breakpoint(body):
return Reply(
status=400,
body=json.dumps(
{
"error": {
"message": "prompt_cache_breakpoint is not supported on this model",
"type": "invalid_request_error",
"param": None,
"code": None,
}
}
).encode(),
)
marker: Final = _request_marker(request)
stream: Final = body.get("stream") is True
response: Final = _responses_body(marker)
if not stream:
return Reply(body=json.dumps(response).encode())
created: Final = {
"type": "response.created",
"sequence_number": 0,
"response": {**response, "status": "in_progress", "output": []},
}
delta: Final = {
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": f"msg_{marker}",
"output_index": 0,
"content_index": 0,
"delta": f"answer marker-{marker}",
}
completed: Final = {"type": "response.completed", "sequence_number": 2, "response": response}
events: Final = (created, delta, completed)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
)
def _prompt(marker: str, label: str) -> str:
return f"{label} marker-{marker}"
def _simple_chat_body(
model: str,
marker: str,
*,
stream: bool = False,
marked: bool = True,
system_as_string: bool = False,
prompt_cache_options: dict[str, JsonValue] | None = None,
) -> dict[str, JsonValue]:
marker_field: Final = {"prompt_cache_breakpoint": _BREAKPOINT} if marked else {}
user: Final = [{"type": "text", "text": _prompt(marker, "user"), **marker_field}]
messages: Final = (
[{"role": "system", "content": _prompt(marker, "system")}, {"role": "user", "content": user}]
if system_as_string
else [{"role": "user", "content": user}]
)
return _JSON_OBJECT.validate_python(
{
"model": model,
"messages": messages,
"tools": _TOOLS,
"reasoning_effort": "low",
"stream": stream,
"num_retries": 0,
**({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}),
}
)
def _multimodal_chat_body(model: str, marker: str, stream: bool) -> dict[str, JsonValue]:
return _JSON_OBJECT.validate_python(
{
"model": model,
"messages": [
{
"role": "system",
"content": [
{
"type": "text",
"text": _prompt(marker, "system"),
"prompt_cache_breakpoint": _BREAKPOINT,
}
],
},
{
"role": "user",
"content": [
{"type": "text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT},
{
"type": "image_url",
"image_url": {"url": _IMAGE_URL},
"prompt_cache_breakpoint": _BREAKPOINT,
},
{
"type": "file",
"file": {"file_id": "file-abc"},
"prompt_cache_breakpoint": _BREAKPOINT,
},
{"type": "text", "text": "unmarked extra text"},
],
},
],
"tools": _TOOLS,
"reasoning_effort": "low",
"stream": stream,
"num_retries": 0,
"prompt_cache_options": {"mode": "explicit"},
}
)
def _expected_multimodal_input(marker: str) -> list[JsonValue]:
return [
{
"type": "message",
"role": "system",
"content": [
{"type": "input_text", "text": _prompt(marker, "system"), "prompt_cache_breakpoint": _BREAKPOINT}
],
},
{
"type": "message",
"role": "user",
"content": [
{"type": "input_text", "text": _prompt(marker, "user"), "prompt_cache_breakpoint": _BREAKPOINT},
{
"type": "input_image",
"image_url": _IMAGE_URL,
"detail": "auto",
"prompt_cache_breakpoint": _BREAKPOINT,
},
{"type": "input_file", "file_id": "file-abc", "prompt_cache_breakpoint": _BREAKPOINT},
{"type": "input_text", "text": "unmarked extra text"},
],
},
]
def _simple_expected_input(marker: str, *, marked: bool) -> list[JsonValue]:
text_block: Final = {"type": "input_text", "text": _prompt(marker, "user")}
return [
{
"type": "message",
"role": "user",
"content": [{**text_block, **({"prompt_cache_breakpoint": _BREAKPOINT} if marked else {})}],
}
]
def _request_body(request: Request) -> dict[str, JsonValue]:
assert request.method == "POST" and request.target == "/v1/responses", request.target
return _JSON_OBJECT.validate_json(request.body)
def _decoded_response_id(response_id: str) -> str:
decoded: Final = _RU._decode_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # reuse ID decoder
response_id
)
raw_response_id: Final = decoded.get("response_id")
assert isinstance(raw_response_id, str), decoded
return raw_response_id
def _spend_request_id_matches(
row: Mapping[str, JsonValue],
caller_response_id: str,
peer_response_id: str,
surface: _Surface,
) -> bool:
request_id: Final = row.get("request_id")
if not isinstance(request_id, str):
return False
match surface:
case "responses":
return request_id == caller_response_id
case "chat":
return _decoded_response_id(request_id) == peer_response_id
def _spend_rows(
model: str,
caller_response_id: str,
peer_response_id: str,
surface: _Surface,
) -> tuple[dict[str, JsonValue], ...]:
def matching_rows(rows: list[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]:
return tuple(
row for row in rows if _spend_request_id_matches(row, caller_response_id, peer_response_id, surface)
)
rows: Final = eventually(
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda candidates: len(matching_rows(candidates)) == 1,
seconds=60,
)
matched: Final = matching_rows(rows)
assert len(matched) == 1, matched
return matched
def _response_id_from_chat_stream(text: str) -> str:
payloads: Final = tuple(
_JSON_OBJECT.validate_json(line.removeprefix("data: "))
for line in text.splitlines()
if line.startswith("data: {")
)
assert payloads, text
response_id: Final = payloads[0].get("id")
assert isinstance(response_id, str), payloads[0]
return response_id
def _extra_body(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
return {
key: value
for key, value in body.items()
if key not in {"model", "messages", "tools", "reasoning_effort", "stream", "num_retries"}
}
def _sync_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served:
base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1"
model: Final = str(body["model"])
messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"])
tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"])
extras: Final = _extra_body(body)
with OpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client:
if stream:
response_stream: Final = client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=True,
extra_body=extras,
)
chunks: Final = tuple(response_stream)
assert chunks
return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "")
response: Final = client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=False,
extra_body=extras,
)
return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "")
async def _async_sdk_chat(gateway: Gateway, body: dict[str, JsonValue], stream: bool) -> _Served:
base_url: Final = f"{str(gateway.client.base_url).rstrip('/')}/v1"
model: Final = str(body["model"])
messages: Final = cast(Iterable[ChatCompletionMessageParam], body["messages"])
tools: Final = cast(Iterable[ChatCompletionToolUnionParam], body["tools"])
extras: Final = _extra_body(body)
async with AsyncOpenAI(api_key=gateway.key, base_url=base_url, max_retries=0) as client:
if stream:
response_stream: Final = await client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=True,
extra_body=extras,
)
chunks: Final = tuple([chunk async for chunk in response_stream])
assert chunks
return _Served(_Call("chat", True, _request_marker_from_body(body)), 200, chunks[0].id, "")
response: Final = await client.chat.completions.create(
model=model,
messages=messages,
tools=tools,
reasoning_effort="low",
stream=False,
extra_body=extras,
)
return _Served(_Call("chat", False, _request_marker_from_body(body)), 200, response.id, "")
async def _serve_chat(
gateway: Gateway,
body: dict[str, JsonValue],
client_kind: _ClientKind,
stream: bool,
) -> _Served:
match client_kind:
case "openai_sync":
return _sync_sdk_chat(gateway, body, stream)
case "openai_async":
return await _async_sdk_chat(gateway, body, stream)
case "httpx":
async with httpx.AsyncClient(
base_url=str(gateway.client.base_url),
headers={"Authorization": f"Bearer {gateway.key}"},
timeout=20,
trust_env=False,
) as client:
return await _raw_call(
client,
"/v1/chat/completions",
body,
_Call("chat", stream, _request_marker_from_body(body)),
)
def _request_marker_from_body(body: Mapping[str, JsonValue]) -> str:
match: Final = _MARKER.search(json.dumps(body).encode())
assert match is not None, body
return match.group(1).decode()
async def _raw_call(
client: httpx.AsyncClient,
path: str,
body: Mapping[str, JsonValue],
call: _Call,
) -> _Served:
async with client.stream(
"POST",
path,
json=body,
headers={"Authorization": f"Bearer {client.headers['Authorization'].removeprefix('Bearer ')}"},
) as response:
content: Final = await response.aread()
status: Final = response.status_code
text: Final = content.decode()
response_id: Final = (
_response_id_from_chat_stream(text)
if status == 200 and call.surface == "chat" and call.stream
else _JSON_OBJECT.validate_json(content).get("id")
if status == 200
else None
)
return _Served(call, status, response_id if isinstance(response_id, str) else None, text)
async def _send_call(
client: httpx.AsyncClient,
model: str,
call: _Call,
) -> _Served:
body: Final = (
_simple_chat_body(model, call.marker, stream=call.stream)
if call.surface == "chat"
else {
"model": model,
"input": _simple_expected_input(call.marker, marked=True),
"stream": call.stream,
"num_retries": 0,
}
)
path: Final = "/v1/chat/completions" if call.surface == "chat" else "/v1/responses"
try:
return await _raw_call(client, path, body, call)
except httpx.TransportError as error:
return _Served(call, 0, None, f"{type(error).__name__}: {error}")
async def _burst(
base_url: str,
key: str,
model: str,
calls: tuple[_Call, ...],
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(
base_url=base_url,
headers={"Authorization": f"Bearer {key}"},
timeout=20,
trust_env=False,
limits=httpx.Limits(max_connections=100),
) as client:
return tuple(await asyncio.gather(*(_send_call(client, model, call) for call in calls)))
def _calls(count: int) -> tuple[_Call, ...]:
surfaces: Final[tuple[_Surface, ...]] = ("chat", "chat", "responses")
return tuple(
_Call(
surface=surfaces[index % len(surfaces)],
stream=index % 3 == 1,
marker=uuid.uuid4().hex,
)
for index in range(count)
)
def _requests_for_marker(requests: tuple[Request, ...], marker: str) -> tuple[Request, ...]:
return tuple(
request
for request in requests
if request.method == "POST" and request.target == "/v1/responses" and _request_marker(request) == marker
)
def _peer_request_has_marker(request: Request, marker: str) -> bool:
body: Final = _request_body(request)
return _contains_breakpoint(body) and _request_marker(request) == marker
def _peer_marker_matches_response(served: _Served, requests: tuple[Request, ...]) -> bool:
peer_requests: Final = _requests_for_marker(requests, served.call.marker)
assert len(peer_requests) == 1, (served, peer_requests)
(peer_request,) = peer_requests
return _peer_request_has_marker(peer_request, served.call.marker)
def _assert_spend_for_result(served: _Served, model: str) -> None:
assert served.status == 200 and served.response_id is not None, served
peer_response_id: Final = _response_id(served.call.marker)
match served.call.surface:
case "responses":
(row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface)
case "chat":
assert _decoded_response_id(served.response_id) == peer_response_id, served
(row,) = _spend_rows(model, served.response_id, peer_response_id, served.call.surface)
request_id: Final = row.get("request_id")
assert isinstance(request_id, str), row
match served.call.surface:
case "responses":
assert request_id == served.response_id, row
case "chat":
assert _decoded_response_id(request_id) == served.response_id, row
from pydantic import JsonValue
from integration._support.client import Gateway
from integration._support.wire import wire_server
from integration.providers._responses_bridge_prompt_cache_breakpoint import (
_BREAKPOINT,
_Call,
_ClientKind,
_JSON_OBJECT,
_MODEL,
_UNSUPPORTED_MODEL,
_assert_spend_for_result,
_contains_breakpoint,
_expected_multimodal_input,
_multimodal_chat_body,
_prompt,
_raw_call,
_request_body,
_responses_reply,
_serve_chat,
_simple_chat_body,
_simple_expected_input,
)
@pytest.mark.parametrize("client_kind", ("openai_sync", "openai_async", "httpx"))
@pytest.mark.parametrize("stream", (False, True))
@ -559,7 +48,6 @@ async def test_caller_prompt_cache_breakpoints_survive_chat_to_responses_bridge(
assert peer_body["input"] == _expected_multimodal_input(marker), peer_body
assert peer_body["prompt_cache_options"] == {"mode": "explicit"}, peer_body
def _expected_uninjected_system_bridge_body(
marker: str,
prompt_cache_options: dict[str, JsonValue] | None = None,
@ -582,7 +70,6 @@ def _expected_uninjected_system_bridge_body(
**({"prompt_cache_options": prompt_cache_options} if prompt_cache_options is not None else {}),
}
async def test_deployment_cache_control_injection_without_options_is_unchanged(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_responses_reply) as wire, gateway.scenario() as scenario:
@ -607,7 +94,6 @@ async def test_deployment_cache_control_injection_without_options_is_unchanged(g
assert not _contains_breakpoint(peer_body), peer_body
assert "prompt_cache_options" not in peer_body, peer_body
async def test_deployment_prompt_cache_options_override_is_unchanged(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
options: Final[dict[str, JsonValue]] = {"mode": "implicit", "ttl": "30m"}
@ -632,7 +118,6 @@ async def test_deployment_prompt_cache_options_override_is_unchanged(gateway: Ga
assert peer_body == _expected_uninjected_system_bridge_body(marker, options), peer_body
assert not _contains_breakpoint(peer_body), peer_body
async def test_unmarked_bridge_and_direct_responses_marker_are_forwarded_unchanged(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_responses_reply) as wire, gateway.scenario() as scenario:
@ -689,7 +174,6 @@ async def test_unmarked_bridge_and_direct_responses_marker_are_forwarded_unchang
(direct_peer,) = wire.drain()
assert _request_body(direct_peer)["input"] == direct_input, _request_body(direct_peer)
@pytest.mark.parametrize("stream", (False, True))
async def test_unsupported_model_drops_breakpoints_without_rejecting_the_request(
gateway: Gateway,
@ -715,197 +199,3 @@ async def test_unsupported_model_drops_breakpoints_without_rejecting_the_request
peer_body: Final = _request_body(peer_request)
assert peer_body["input"] == _simple_expected_input(marker, marked=False), peer_body
assert not _contains_breakpoint(peer_body), peer_body
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
base_config: Final = _JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
config: Final = {
**base_config,
"model_list": [
{
"model_name": _CONFIG_MODEL,
"litellm_params": {
"model": _MODEL,
"api_base": wire.url + "/v1",
"api_key": _API_KEY,
},
},
],
}
path: Final = tmp_path / "responses-bridge-cache-breakpoint-chaos.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _open_upstream_connections(pid: int, port: int) -> int:
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(180)
async def test_worker_and_peer_outages_preserve_markers_and_recover(
gateway: Gateway,
tmp_path: Path,
) -> None:
calls: Final = _calls(30)
release: Final = threading.Event()
early_release: Final = threading.Event()
outage_release: Final = threading.Event()
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
early_calls: Final = calls[:10]
early_markers: Final = frozenset(call.marker for call in early_calls)
def held(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return _responses_reply(request)
marker: Final = _request_marker(request)
held_markers.put(marker)
gate: Final = early_release if marker in early_markers else release
assert gate.wait(timeout=60), "The worker-kill burst was never released"
return _responses_reply(request)
with ExitStack() as peer_stack:
wire: Final = peer_stack.enter_context(wire_server(held))
config: Final = _chaos_config(wire, tmp_path)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
try:
candidate: Final = owned.gateway
workers: Final[tuple[int, ...]] = eventually(
lambda: tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
async with httpx.AsyncClient(
base_url=str(candidate.client.base_url),
headers={"Authorization": f"Bearer {candidate.key}"},
timeout=20,
trust_env=False,
limits=httpx.Limits(max_connections=100),
) as client:
burst_tasks: Final = tuple(
asyncio.create_task(_send_call(client, _CONFIG_MODEL, call)) for call in calls
)
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == len(calls), 60)
early_release.set()
early_served: Final = await asyncio.gather(*burst_tasks[: len(early_calls)])
early_successful: Final = tuple(item for item in early_served if item.status == 200)
for item in early_successful:
_assert_spend_for_result(item, _CONFIG_MODEL)
upstream_port_value: Final = urlsplit(wire.url).port
assert upstream_port_value is not None
upstream_port: Final = upstream_port_value
active_by_worker: Final = eventually(
lambda: {pid: _open_upstream_connections(pid, upstream_port) for pid in workers},
lambda counts: sum(counts.values()) == len(calls) - len(early_calls),
seconds=30,
)
victim_pid: Final = max(workers, key=active_by_worker.__getitem__)
survivor_pids: Final = tuple(pid for pid in workers if pid != victim_pid)
assert active_by_worker[victim_pid] > 0 and len(survivor_pids) == 1, active_by_worker
(survivor_pid,) = survivor_pids
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
remaining_served: Final = await asyncio.gather(*burst_tasks[len(early_calls) :])
served: Final = (*early_served, *remaining_served)
successful: Final = tuple(item for item in served if item.status == 200)
connection_errors: Final = tuple(item for item in served if item.status == 0)
print(f"worker-kill burst: {len(successful)} HTTP 200, {len(connection_errors)} connection errors")
assert len(successful) + len(connection_errors) == len(calls), {
"successes": len(successful),
"connection_errors": len(connection_errors),
"responses": served,
}
assert successful and connection_errors, {
"successes": len(successful),
"connection_errors": len(connection_errors),
}
follow_ups: Final = (
_Call("chat", False, uuid.uuid4().hex),
_Call("responses", False, uuid.uuid4().hex),
)
recovered: Final = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
follow_ups,
)
assert all(item.status == 200 for item in recovered), recovered
assert psutil.pid_exists(survivor_pid), survivor_pid
received_after_worker_kill: Final = wire.drain()
worker_marker_failures: Final = tuple(
item.call.marker
for item in (*successful, *recovered)
if not _peer_marker_matches_response(item, received_after_worker_kill)
)
for item in (*successful, *recovered):
assert item.response_id is not None, item
_assert_spend_for_result(item, _CONFIG_MODEL)
peer_stack.close()
outage_seen: Final[SimpleQueue[str]] = SimpleQueue()
def outage(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return _responses_reply(request)
outage_seen.put(_request_marker(request))
assert outage_release.wait(timeout=60), "The peer-outage burst was never stopped"
return Reply(
status=503,
body=b'{"error":{"message":"synthetic peer outage","type":"server_error"}}',
)
peer_stack.enter_context(wire_server(outage, port=upstream_port))
outage_calls: Final = _calls(12)
outage_burst: Final = asyncio.create_task(
_burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, outage_calls)
)
await asyncio.to_thread(eventually, outage_seen.qsize, lambda size: size == len(outage_calls), 30)
outage_release.set()
peer_stack.close()
outage_served: Final = await outage_burst
assert len(outage_served) == len(outage_calls), outage_served
assert all(item.status >= 400 and "error" in item.text.lower() for item in outage_served), outage_served
down_call: Final = _Call("chat", False, uuid.uuid4().hex)
(down_response,) = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
(down_call,),
)
assert down_response.status >= 400 and "error" in down_response.text.lower(), down_response
restarted_wire: Final = peer_stack.enter_context(wire_server(_responses_reply, port=upstream_port))
recovery_calls: Final = (
_Call("chat", False, uuid.uuid4().hex),
_Call("responses", False, uuid.uuid4().hex),
)
recovered_after_peer_restart: Final = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
recovery_calls,
)
assert all(item.status == 200 for item in recovered_after_peer_restart), recovered_after_peer_restart
restarted_requests: Final = restarted_wire.drain()
recovery_marker_failures: Final = tuple(
item.call.marker
for item in recovered_after_peer_restart
if not _peer_marker_matches_response(item, restarted_requests)
)
for item in recovered_after_peer_restart:
assert item.response_id is not None, item
_assert_spend_for_result(item, _CONFIG_MODEL)
assert not (*worker_marker_failures, *recovery_marker_failures), {
"worker_marker_failures": worker_marker_failures,
"recovery_marker_failures": recovery_marker_failures,
}
finally:
release.set()
outage_release.set()

View file

@ -0,0 +1,229 @@
from __future__ import annotations
import asyncio
import re
import signal
import threading
import uuid
from contextlib import ExitStack
from pathlib import Path
from queue import SimpleQueue
from typing import Final
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._responses_bridge_prompt_cache_breakpoint import (
_Call,
_JSON_OBJECT,
_MODEL,
_assert_spend_for_result,
_burst,
_calls,
_peer_marker_matches_response,
_request_marker,
_responses_reply,
_send_call,
)
_CONFIG_MODEL: Final = "responses-bridge-cache-breakpoint-chaos"
_API_KEY: Final = "synthetic-responses-bridge-key"
_STARTED_WORKER: Final[re.Pattern[str]] = re.compile(r"Started server process \[(\d+)\]")
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
base_config: Final = _JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
config: Final = {
**base_config,
"model_list": [
{
"model_name": _CONFIG_MODEL,
"litellm_params": {
"model": _MODEL,
"api_base": wire.url + "/v1",
"api_key": _API_KEY,
},
},
],
}
path: Final = tmp_path / "responses-bridge-cache-breakpoint-chaos.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _open_upstream_connections(pid: int, port: int) -> int:
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(180)
async def test_worker_and_peer_outages_preserve_markers_and_recover(
gateway: Gateway,
tmp_path: Path,
) -> None:
calls: Final = _calls(30)
release: Final = threading.Event()
early_release: Final = threading.Event()
outage_release: Final = threading.Event()
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
early_calls: Final = calls[:10]
early_markers: Final = frozenset(call.marker for call in early_calls)
def held(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return _responses_reply(request)
marker: Final = _request_marker(request)
held_markers.put(marker)
gate: Final = early_release if marker in early_markers else release
assert gate.wait(timeout=60), "The worker-kill burst was never released"
return _responses_reply(request)
with ExitStack() as peer_stack:
wire: Final = peer_stack.enter_context(wire_server(held))
config: Final = _chaos_config(wire, tmp_path)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
try:
candidate: Final = owned.gateway
workers: Final[tuple[int, ...]] = eventually(
lambda: tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
async with httpx.AsyncClient(
base_url=str(candidate.client.base_url),
headers={"Authorization": f"Bearer {candidate.key}"},
timeout=20,
trust_env=False,
limits=httpx.Limits(max_connections=100),
) as client:
burst_tasks: Final = tuple(
asyncio.create_task(_send_call(client, _CONFIG_MODEL, call)) for call in calls
)
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == len(calls), 60)
early_release.set()
early_served: Final = await asyncio.gather(*burst_tasks[: len(early_calls)])
early_successful: Final = tuple(item for item in early_served if item.status == 200)
for item in early_successful:
_assert_spend_for_result(item, _CONFIG_MODEL)
upstream_port_value: Final = urlsplit(wire.url).port
assert upstream_port_value is not None
upstream_port: Final = upstream_port_value
active_by_worker: Final = eventually(
lambda: {pid: _open_upstream_connections(pid, upstream_port) for pid in workers},
lambda counts: sum(counts.values()) == len(calls) - len(early_calls),
seconds=30,
)
victim_pid: Final = max(workers, key=active_by_worker.__getitem__)
survivor_pids: Final = tuple(pid for pid in workers if pid != victim_pid)
assert active_by_worker[victim_pid] > 0 and len(survivor_pids) == 1, active_by_worker
(survivor_pid,) = survivor_pids
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
remaining_served: Final = await asyncio.gather(*burst_tasks[len(early_calls) :])
served: Final = (*early_served, *remaining_served)
successful: Final = tuple(item for item in served if item.status == 200)
connection_errors: Final = tuple(item for item in served if item.status == 0)
print(f"worker-kill burst: {len(successful)} HTTP 200, {len(connection_errors)} connection errors")
assert len(successful) + len(connection_errors) == len(calls), {
"successes": len(successful),
"connection_errors": len(connection_errors),
"responses": served,
}
assert successful and connection_errors, {
"successes": len(successful),
"connection_errors": len(connection_errors),
}
follow_ups: Final = (
_Call("chat", False, uuid.uuid4().hex),
_Call("responses", False, uuid.uuid4().hex),
)
recovered: Final = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
follow_ups,
)
assert all(item.status == 200 for item in recovered), recovered
assert psutil.pid_exists(survivor_pid), survivor_pid
received_after_worker_kill: Final = wire.drain()
worker_marker_failures: Final = tuple(
item.call.marker
for item in (*successful, *recovered)
if not _peer_marker_matches_response(item, received_after_worker_kill)
)
for item in (*successful, *recovered):
assert item.response_id is not None, item
_assert_spend_for_result(item, _CONFIG_MODEL)
peer_stack.close()
outage_seen: Final[SimpleQueue[str]] = SimpleQueue()
def outage(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return _responses_reply(request)
outage_seen.put(_request_marker(request))
assert outage_release.wait(timeout=60), "The peer-outage burst was never stopped"
return Reply(
status=503,
body=b'{"error":{"message":"synthetic peer outage","type":"server_error"}}',
)
peer_stack.enter_context(wire_server(outage, port=upstream_port))
outage_calls: Final = _calls(12)
outage_burst: Final = asyncio.create_task(
_burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, outage_calls)
)
await asyncio.to_thread(eventually, outage_seen.qsize, lambda size: size == len(outage_calls), 30)
outage_release.set()
peer_stack.close()
outage_served: Final = await outage_burst
assert len(outage_served) == len(outage_calls), outage_served
assert all(item.status >= 400 and "error" in item.text.lower() for item in outage_served), outage_served
down_call: Final = _Call("chat", False, uuid.uuid4().hex)
(down_response,) = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
(down_call,),
)
assert down_response.status >= 400 and "error" in down_response.text.lower(), down_response
restarted_wire: Final = peer_stack.enter_context(wire_server(_responses_reply, port=upstream_port))
recovery_calls: Final = (
_Call("chat", False, uuid.uuid4().hex),
_Call("responses", False, uuid.uuid4().hex),
)
recovered_after_peer_restart: Final = await _burst(
str(candidate.client.base_url),
candidate.key,
_CONFIG_MODEL,
recovery_calls,
)
assert all(item.status == 200 for item in recovered_after_peer_restart), recovered_after_peer_restart
restarted_requests: Final = restarted_wire.drain()
recovery_marker_failures: Final = tuple(
item.call.marker
for item in recovered_after_peer_restart
if not _peer_marker_matches_response(item, restarted_requests)
)
for item in recovered_after_peer_restart:
assert item.response_id is not None, item
_assert_spend_for_result(item, _CONFIG_MODEL)
assert not (*worker_marker_failures, *recovery_marker_failures), {
"worker_marker_failures": worker_marker_failures,
"recovery_marker_failures": recovery_marker_failures,
}
finally:
release.set()
outage_release.set()

View file

@ -4291,6 +4291,42 @@ def test_prompt_cache_breakpoints_are_dropped_for_unsupported_models() -> None:
]
@pytest.mark.parametrize(
("litellm_params", "keep_marker"),
(({"base_model": "gpt-5.6"}, True), ({}, False)),
ids=("supported-base-model", "missing-base-model"),
)
def test_prompt_cache_breakpoint_supports_model_alias_with_base_model(
litellm_params: dict[str, object],
keep_marker: bool,
) -> None:
handler: Final = LiteLLMResponsesTransformationHandler()
cache_breakpoint: Final = {"mode": "explicit"}
marked_content: Final = {"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": cache_breakpoint}
request: Final = handler.transform_request(
model="mydeployment",
messages=[{"role": "user", "content": [marked_content]}],
optional_params={},
litellm_params=litellm_params,
headers={},
litellm_logging_obj=Mock(),
)
expected_content: Final = {
"type": "input_text",
"text": "Stable prefix",
**({"prompt_cache_breakpoint": cache_breakpoint} if keep_marker else {}),
}
assert request["input"] == [
{
"type": "message",
"role": "user",
"content": [expected_content],
}
]
def test_mid_conversation_system_string_stays_in_input_after_a_user_turn():
handler: Final = LiteLLMResponsesTransformationHandler()