mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(responses): honor caller stream flag when provider forces SSE (internal copy of #34095) (#41235)
* fix(responses): honor caller stream flag when provider forces SSE The Responses handlers decided whether to hand back a streaming iterator from the provider payload's `stream` field, which chatgpt sets unconditionally because the Codex backend only serves SSE. A caller that sent `stream: false` therefore received a raw SSE stream on /v1/responses, and the chat-completions bridge failed with "Unknown items in responses API response: []" once its recovery path lost the raw SSE it reads from Transport streaming still follows the provider payload; only the caller's own `stream` value now decides the response shape. When the provider forces SSE for a non-streaming caller the body is read and aggregated through the existing path * test(chatgpt): inject the authenticator into the responses config so handler tests never log in * fix(responses): treat an extra_body stream flag as the caller's own and drop a redundant comment * test(integration): audit the chatgpt caller stream flag across responses, chat and messages --------- Co-authored-by: SeongWoon Cho <coffee@soylatte.kr> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
e366e72502
commit
a983b2a5e7
7 changed files with 1285 additions and 31 deletions
|
|
@ -39,9 +39,9 @@ _CHATGPT_SERVICE_TIERS: Final = {"default": "default", "priority": "priority", "
|
|||
|
||||
|
||||
class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, authenticator: Authenticator | None = None) -> None:
|
||||
super().__init__()
|
||||
self.authenticator = Authenticator()
|
||||
self.authenticator = authenticator if authenticator is not None else Authenticator()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
|
|
|
|||
|
|
@ -2507,6 +2507,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
# Check if streaming is requested
|
||||
stream = response_api_optional_request_params.get("stream", False)
|
||||
caller_requested_stream: Final = bool(stream or (extra_body or {}).get("stream"))
|
||||
|
||||
api_base: Final = responses_api_provider_config.get_complete_url(
|
||||
api_base=litellm_params.api_base,
|
||||
|
|
@ -2585,8 +2586,20 @@ class BaseLLMHTTPHandler:
|
|||
stream=stream,
|
||||
**body_kwargs,
|
||||
)
|
||||
if fake_stream is True:
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
if caller_requested_stream:
|
||||
if fake_stream is True:
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
return SyncResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -2596,17 +2609,7 @@ class BaseLLMHTTPHandler:
|
|||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
return SyncResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
response.read()
|
||||
else:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
|
|
@ -2699,6 +2702,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
# Check if streaming is requested
|
||||
stream = response_api_optional_request_params.get("stream", False)
|
||||
caller_requested_stream: Final = bool(stream or (extra_body or {}).get("stream"))
|
||||
|
||||
api_base: Final = responses_api_provider_config.get_complete_url(
|
||||
api_base=litellm_params.api_base,
|
||||
|
|
@ -2778,8 +2782,20 @@ class BaseLLMHTTPHandler:
|
|||
**body_kwargs,
|
||||
)
|
||||
|
||||
if fake_stream is True:
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
if caller_requested_stream:
|
||||
if fake_stream is True:
|
||||
return MockResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
return ResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -2789,18 +2805,7 @@ class BaseLLMHTTPHandler:
|
|||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
# Return the streaming iterator
|
||||
return ResponsesAPIStreamingIterator(
|
||||
response=response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
await response.aread()
|
||||
else:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
|
|
|
|||
50
tests/integration/_support/agentic_probe.py
Normal file
50
tests/integration/_support/agentic_probe.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from integration._support import responses_vendor as rv
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
OUT_ENVIRONMENT: Final = "AGENTIC_PROBE_OUT"
|
||||
|
||||
|
||||
class AgenticProbe(CustomLogger):
|
||||
"""Records every agentic-loop hook call the proxy makes, one JSON line per call, and never runs a loop."""
|
||||
|
||||
async def async_should_run_agentic_loop(
|
||||
self,
|
||||
response: object,
|
||||
model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
tools: Sequence[Mapping[str, object]] | None,
|
||||
stream: bool,
|
||||
custom_llm_provider: str,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> tuple[bool, dict[str, object]]: # mutable-ok: the CustomLogger hook contract returns a dict
|
||||
out: Final = os.environ.get(OUT_ENVIRONMENT)
|
||||
if out:
|
||||
line: Final[Mapping[str, JsonValue]] = {
|
||||
"marker": rv.newest_marker(f"{messages!s} {response!s}"),
|
||||
"surface": str(kwargs.get("_agentic_loop_api_surface")),
|
||||
"response_type": type(response).__name__,
|
||||
"stream": stream,
|
||||
"model": model,
|
||||
"provider": custom_llm_provider,
|
||||
}
|
||||
with open(out, "a", encoding="utf-8") as sink:
|
||||
sink.write(json.dumps(line) + "\n")
|
||||
return False, {}
|
||||
|
||||
|
||||
probe: Final = AgenticProbe()
|
||||
|
||||
|
||||
def lines(path: Path, marker: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
recorded: Final = tuple(rv.JSON_OBJECT.validate_json(line) for line in path.read_text().splitlines() if line)
|
||||
return tuple(line for line in recorded if line["marker"] == marker)
|
||||
257
tests/integration/_support/codex_vendor.py
Normal file
257
tests/integration/_support/codex_vendor.py
Normal file
|
|
@ -0,0 +1,257 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import yaml
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.wire import Reply, Request
|
||||
from pydantic import JsonValue
|
||||
|
||||
TOKEN: Final = "synthetic-chatgpt-token"
|
||||
ACCOUNT: Final = "acct-synthetic"
|
||||
PROBE_CALLBACK: Final = "integration._support.agentic_probe.probe"
|
||||
INPUT_MUST_BE_A_LIST: Final = "Input must be a list"
|
||||
UNAUTHORIZED: Final = "Unauthorized"
|
||||
UNAUTHORIZED_DIRECTIVE: Final = "codex-unauthorized"
|
||||
FAILED_DIRECTIVE: Final = "codex-failed"
|
||||
INCOMPLETE_DIRECTIVE: Final = "codex-incomplete"
|
||||
INCOMPLETE_PAUSE_SECONDS: Final = 0.5
|
||||
INCOMPLETE_CHUNKS: Final = 4
|
||||
USAGE: Final[Mapping[str, JsonValue]] = {
|
||||
"input_tokens": 30,
|
||||
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0},
|
||||
"output_tokens": 5,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
"total_tokens": 35,
|
||||
}
|
||||
_TOTAL_KEYS: Final = ("input_tokens", "output_tokens", "total_tokens")
|
||||
FORWARDED_KEYS: Final = frozenset(
|
||||
{
|
||||
"model",
|
||||
"input",
|
||||
"instructions",
|
||||
"stream",
|
||||
"store",
|
||||
"include",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"reasoning",
|
||||
"previous_response_id",
|
||||
"truncation",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def login(directory: Path) -> Path:
|
||||
chatgpt: Final = directory / "chatgpt"
|
||||
chatgpt.mkdir()
|
||||
(chatgpt / "auth.json").write_text(
|
||||
json.dumps({"access_token": TOKEN, "account_id": ACCOUNT, "expires_at": time.time() + 3600})
|
||||
)
|
||||
return chatgpt
|
||||
|
||||
|
||||
def proxy_config(directory: Path, *, probe: bool) -> Path:
|
||||
stock: Final = rv.JSON_OBJECT.validate_python(
|
||||
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
)
|
||||
litellm_settings: Final = rv.JSON_OBJECT.validate_python(stock["litellm_settings"])
|
||||
router_settings: Final = rv.JSON_OBJECT.validate_python(stock.get("router_settings") or {})
|
||||
config: Final[Mapping[str, JsonValue]] = {
|
||||
**stock,
|
||||
"litellm_settings": {**litellm_settings, "callbacks": [PROBE_CALLBACK]} if probe else litellm_settings,
|
||||
"router_settings": {**router_settings, "num_retries": 0},
|
||||
}
|
||||
path: Final = directory / "chatgpt-codex-rig.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def totals(usage: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
|
||||
return {key: usage[key] for key in _TOTAL_KEYS}
|
||||
|
||||
|
||||
def failure_message(marker: str | None) -> str:
|
||||
return f"codex failed marker-{marker}"
|
||||
|
||||
|
||||
def forwarded(request: Request, marker: str, *, stream: bool = True) -> Mapping[str, JsonValue]:
|
||||
assert request.method == "POST", request.method
|
||||
assert urlsplit(request.target).path.endswith("/responses"), request.target
|
||||
assert request.headers.get("authorization") == f"Bearer {TOKEN}", request.headers
|
||||
assert request.headers.get("chatgpt-account-id") == ACCOUNT, request.headers
|
||||
body: Final = rv.JSON_OBJECT.validate_json(request.body)
|
||||
assert set(body) <= FORWARDED_KEYS, sorted(body)
|
||||
assert (body["stream"], body["store"]) == (stream, False), body
|
||||
assert isinstance(body["instructions"], str) and body["instructions"], body
|
||||
assert rv.newest_marker(json.dumps(body["input"])) == marker, body["input"]
|
||||
return body
|
||||
|
||||
|
||||
def _detail(status: int, detail: str) -> Reply:
|
||||
return Reply(status=status, body=json.dumps({"detail": detail}).encode())
|
||||
|
||||
|
||||
def _stream(
|
||||
frames: Sequence[Mapping[str, JsonValue]],
|
||||
*,
|
||||
abort_after: int | None = None,
|
||||
pause: float = 0,
|
||||
gate: threading.Event | None = None,
|
||||
) -> Reply:
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(rv.sse(frame) for frame in frames),
|
||||
abort_after=abort_after,
|
||||
pause_between_chunks=pause,
|
||||
gate_after_first=gate,
|
||||
)
|
||||
|
||||
|
||||
def _response(model: str, tag: str) -> Mapping[str, JsonValue]:
|
||||
return {
|
||||
"id": f"resp_{tag}",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "in_progress",
|
||||
"model": model,
|
||||
"output": [],
|
||||
"instructions": "You are a coding agent.",
|
||||
"metadata": {},
|
||||
"parallel_tool_calls": True,
|
||||
"temperature": 1.0,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": 1.0,
|
||||
"reasoning": {"effort": "medium", "summary": None},
|
||||
"text": {"format": {"type": "text"}, "verbosity": "medium"},
|
||||
"truncation": "disabled",
|
||||
"store": False,
|
||||
"background": False,
|
||||
"service_tier": "default",
|
||||
}
|
||||
|
||||
|
||||
def _message(tag: str, text: str) -> Mapping[str, JsonValue]:
|
||||
return {
|
||||
"id": f"msg_{tag}",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"phase": "final_answer",
|
||||
"content": [{"type": "output_text", "annotations": [], "logprobs": [], "text": text}],
|
||||
}
|
||||
|
||||
|
||||
def _frames(model: str, tag: str, text: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
response: Final = _response(model, tag)
|
||||
message: Final = _message(tag, text)
|
||||
part: Final[Mapping[str, JsonValue]] = {"type": "output_text", "annotations": [], "logprobs": [], "text": ""}
|
||||
position: Final[Mapping[str, JsonValue]] = {"item_id": f"msg_{tag}", "output_index": 0, "content_index": 0}
|
||||
return (
|
||||
{"type": "response.created", "sequence_number": 0, "model": model, "response": dict(response)},
|
||||
{"type": "response.in_progress", "sequence_number": 1, "model": model, "response": dict(response)},
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"sequence_number": 2,
|
||||
"output_index": 0,
|
||||
"model": model,
|
||||
"item": {**message, "status": "in_progress", "content": []},
|
||||
},
|
||||
{"type": "response.content_part.added", "sequence_number": 3, "model": model, **position, "part": dict(part)},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 4,
|
||||
"model": model,
|
||||
**position,
|
||||
"delta": text,
|
||||
"logprobs": [],
|
||||
"obfuscation": "",
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.done",
|
||||
"sequence_number": 5,
|
||||
"model": model,
|
||||
**position,
|
||||
"text": text,
|
||||
"logprobs": [],
|
||||
},
|
||||
{
|
||||
"type": "response.content_part.done",
|
||||
"sequence_number": 6,
|
||||
"model": model,
|
||||
**position,
|
||||
"part": {**part, "text": text},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"sequence_number": 7,
|
||||
"output_index": 0,
|
||||
"model": model,
|
||||
"item": dict(message),
|
||||
},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"sequence_number": 8,
|
||||
"model": model,
|
||||
"response": {**response, "status": "completed", "usage": dict(USAGE), "completed_at": 2},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _failed_frames(model: str, tag: str, marker: str | None) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
response: Final = _response(model, tag)
|
||||
return (
|
||||
{"type": "response.created", "sequence_number": 0, "model": model, "response": dict(response)},
|
||||
{
|
||||
"type": "response.failed",
|
||||
"sequence_number": 1,
|
||||
"model": model,
|
||||
"response": {
|
||||
**response,
|
||||
"status": "failed",
|
||||
"error": {"code": "server_error", "message": failure_message(marker)},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CodexVendor:
|
||||
"""The ChatGPT Codex backend as the proxy sees it: SSE only, and `input` must be a list of items."""
|
||||
|
||||
pause_between_chunks: float = 0
|
||||
incomplete_gate: threading.Event | None = None
|
||||
|
||||
def respond(self, request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return Reply(body=json.dumps({"object": "list", "data": [{"id": "gpt-5.5", "object": "model"}]}).encode())
|
||||
assert urlsplit(request.target).path.endswith("/responses"), request.target
|
||||
body: Final = rv.JSON_OBJECT.validate_json(request.body)
|
||||
items: Final = body.get("input")
|
||||
if not isinstance(items, list) or any(not isinstance(item, dict) for item in items):
|
||||
return _detail(400, INPUT_MUST_BE_A_LIST)
|
||||
text: Final = json.dumps(items)
|
||||
marker: Final = rv.newest_marker(text)
|
||||
model: Final = str(body["model"])
|
||||
tag: Final = uuid.uuid4().hex
|
||||
if UNAUTHORIZED_DIRECTIVE in text:
|
||||
return _detail(401, UNAUTHORIZED)
|
||||
if FAILED_DIRECTIVE in text:
|
||||
return _stream(_failed_frames(model, tag, marker))
|
||||
if INCOMPLETE_DIRECTIVE in text:
|
||||
return _stream(
|
||||
_frames(model, tag, rv.answer(marker)),
|
||||
abort_after=INCOMPLETE_CHUNKS,
|
||||
pause=INCOMPLETE_PAUSE_SECONDS,
|
||||
gate=self.incomplete_gate,
|
||||
)
|
||||
return _stream(_frames(model, tag, rv.answer(marker)), pause=self.pause_between_chunks)
|
||||
|
|
@ -0,0 +1,319 @@
|
|||
import asyncio
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import 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 codex_vendor as cv
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(240)
|
||||
|
||||
_MODEL: Final = "chatgpt/gpt-5.5"
|
||||
_CONFIG_MODEL: Final = "chatgpt-caller-stream-chaos"
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_NO_CACHE: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}}
|
||||
_ENDPOINTS: Final = ("responses", "chat", "messages")
|
||||
|
||||
Endpoint: TypeAlias = Literal["responses", "chat", "messages"]
|
||||
|
||||
|
||||
@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 _Rig:
|
||||
port: int
|
||||
proxy: OwnedProxy
|
||||
|
||||
@property
|
||||
def gateway(self) -> Gateway:
|
||||
return self.proxy.gateway
|
||||
|
||||
@property
|
||||
def vendor_url(self) -> str:
|
||||
return f"http://127.0.0.1:{self.port}"
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with wire_server(cv.CodexVendor().respond) as probe:
|
||||
port: Final = urlsplit(probe.url).port
|
||||
assert port is not None, probe.url
|
||||
return port
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
|
||||
directory: Final = tmp_path_factory.mktemp("chatgpt-caller-stream-chaos")
|
||||
port: Final = _free_port()
|
||||
overrides: Final = {"CHATGPT_TOKEN_DIR": str(cv.login(directory)), "CHATGPT_API_BASE": f"http://127.0.0.1:{port}"}
|
||||
config: Final = cv.proxy_config(directory, probe=False)
|
||||
with gateway_from_environment() as gateway:
|
||||
with owned_proxy_process(gateway, directory, overrides, config=config, workers=2) as owned:
|
||||
yield _Rig(port, owned)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model(rig: _Rig) -> Iterator[str]:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
yield scenario.model(model=_MODEL, api_base=rig.vendor_url, api_key=None)
|
||||
|
||||
|
||||
def _path(endpoint: Endpoint) -> str:
|
||||
match endpoint:
|
||||
case "responses":
|
||||
return "/v1/responses"
|
||||
case "chat":
|
||||
return "/v1/chat/completions"
|
||||
case "messages":
|
||||
return "/v1/messages"
|
||||
|
||||
|
||||
def _body(model: str, call: _Call) -> Mapping[str, JsonValue]:
|
||||
prompt: Final = f"Say marker-{call.marker}"
|
||||
common: Final[Mapping[str, JsonValue]] = {"model": model, "stream": call.stream, "num_retries": 0, **_NO_CACHE}
|
||||
match call.endpoint:
|
||||
case "responses":
|
||||
return {**common, "input": [{"role": "user", "content": prompt}]}
|
||||
case "chat":
|
||||
return {**common, "messages": [{"role": "user", "content": prompt}]}
|
||||
case "messages":
|
||||
return {**common, "max_tokens": 64, "messages": [{"role": "user", "content": prompt}]}
|
||||
|
||||
|
||||
def _calls(count: int, endpoints: tuple[Endpoint, ...]) -> tuple[_Call, ...]:
|
||||
return tuple(
|
||||
_Call(endpoint=endpoints[index % len(endpoints)], stream=index % 2 == 1, marker=uuid.uuid4().hex)
|
||||
for index in range(count)
|
||||
)
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
_path(call.endpoint),
|
||||
json=_body(model, call),
|
||||
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, ...], *, 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) 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 not served.call.stream:
|
||||
return string_value(rv.JSON_OBJECT.validate_json(served.text)["id"])
|
||||
frames: Final = _frames(served.text)
|
||||
match served.call.endpoint:
|
||||
case "responses":
|
||||
(completed,) = [frame for frame in frames if frame.get("type") == "response.completed"]
|
||||
return string_value(rv.JSON_OBJECT.validate_python(completed["response"])["id"])
|
||||
case "chat":
|
||||
return string_value(frames[0]["id"])
|
||||
case "messages":
|
||||
return None
|
||||
|
||||
|
||||
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")
|
||||
|
||||
|
||||
def _marked(received: tuple[Request, ...]) -> Mapping[str, Request]:
|
||||
posts: Final = tuple(request for request in received if request.method == "POST")
|
||||
marked: Final = {marker: request for request in posts if (marker := rv.newest_marker(request.body.decode()))}
|
||||
assert len(marked) == len(posts), [request.body for request in posts]
|
||||
return marked
|
||||
|
||||
|
||||
def _assert_forwarded(forwarded: Mapping[str, Request], calls: tuple[_Call, ...]) -> None:
|
||||
assert set(forwarded) == {call.marker for call in calls}, sorted(forwarded)
|
||||
for marker, request in forwarded.items():
|
||||
cv.forwarded(request, marker)
|
||||
|
||||
|
||||
def _spend_rows(model: str, expected: int) -> Sequence[Mapping[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, litellm_call_id, status FROM "LiteLLM_SpendLogs" WHERE model_group = %s', (model,)
|
||||
),
|
||||
lambda found: len(found) >= expected,
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def _assert_each_lands_once(
|
||||
rows: Sequence[Mapping[str, JsonValue]], failed: tuple[_Served, ...], served: tuple[_Served, ...]
|
||||
) -> None:
|
||||
by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows}
|
||||
assert len(by_call) == len(rows) == len(failed) + len(served), rows
|
||||
for item in failed:
|
||||
assert by_call[item.call_id]["status"] == "failure", (item.call_id, rows)
|
||||
for item in served:
|
||||
_assert_served_landed(by_call[item.call_id], item)
|
||||
|
||||
|
||||
def _assert_served_landed(row: Mapping[str, JsonValue], item: _Served) -> None:
|
||||
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)
|
||||
|
||||
|
||||
async def test_mixed_burst_across_the_three_endpoints_answers_each_in_its_own_shape(rig: _Rig, model: str) -> None:
|
||||
calls: Final = _calls(24, _ENDPOINTS)
|
||||
with wire_server(cv.CodexVendor().respond, port=rig.port) as wire:
|
||||
served: Final = await _burst(rig.gateway, model, calls)
|
||||
assert len(served) == 24
|
||||
for item in served:
|
||||
_assert_answered_in_its_own_shape(item)
|
||||
_assert_forwarded(_marked(wire.drain()), calls)
|
||||
_assert_each_lands_once(_spend_rows(model, 24), (), served)
|
||||
|
||||
|
||||
async def test_vendor_outage_fails_each_call_cleanly_and_the_restarted_vendor_serves_the_next_burst(
|
||||
rig: _Rig, model: str
|
||||
) -> None:
|
||||
while_down: Final = _calls(12, _ENDPOINTS)
|
||||
after: Final = _calls(12, _ENDPOINTS)
|
||||
failed: Final = await _burst(rig.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
|
||||
assert item.call_id, item
|
||||
down: Final = _health(rig.gateway, model)
|
||||
assert (down["healthy_count"], down["unhealthy_count"]) == (0, 1), down
|
||||
with wire_server(cv.CodexVendor().respond, port=rig.port) as wire:
|
||||
_health(rig.gateway, model)
|
||||
probes: Final = wire.drain()
|
||||
assert [rv.newest_marker(request.body.decode()) for request in probes if request.method == "POST"] == [None]
|
||||
served: Final = await _burst(rig.gateway, model, after)
|
||||
assert len(served) == 12
|
||||
for item in served:
|
||||
_assert_answered_in_its_own_shape(item)
|
||||
_assert_forwarded(_marked(wire.drain()), after)
|
||||
_assert_each_lands_once(_spend_rows(model, 24), failed, served)
|
||||
|
||||
|
||||
def _chaos_config(vendor_url: str, directory: Path) -> Path:
|
||||
config: Final = rv.JSON_OBJECT.validate_python(yaml.safe_load(cv.proxy_config(directory, probe=False).read_text()))
|
||||
path: Final = directory / "chatgpt-caller-stream-worker-chaos.yaml"
|
||||
path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
**config,
|
||||
"model_list": [
|
||||
{"model_name": _CONFIG_MODEL, "litellm_params": {"model": _MODEL, "api_base": vendor_url}}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
return path
|
||||
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_answering_json(gateway: Gateway, tmp_path: Path) -> None:
|
||||
calls: Final = tuple(_Call("responses", False, uuid.uuid4().hex) for _ in range(20))
|
||||
release: Final = threading.Event()
|
||||
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
vendor: Final = cv.CodexVendor()
|
||||
|
||||
def held(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return vendor.respond(request)
|
||||
marker: Final = rv.newest_marker(request.body.decode())
|
||||
assert marker is not None, request.body
|
||||
held_markers.put(marker)
|
||||
assert release.wait(timeout=60), "The burst was never released"
|
||||
return vendor.respond(request)
|
||||
|
||||
with wire_server(held) as wire:
|
||||
overrides: Final = {"CHATGPT_TOKEN_DIR": str(cv.login(tmp_path)), "CHATGPT_API_BASE": wire.url}
|
||||
with owned_proxy_process(
|
||||
gateway, tmp_path, overrides, config=_chaos_config(wire.url, tmp_path), workers=2
|
||||
) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
workers: Final = eventually(
|
||||
lambda: tuple(int(found.group(1)) for found in _STARTED_WORKER.finditer(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
burst: Final = asyncio.create_task(_burst(candidate, _CONFIG_MODEL, calls, tolerate_transport_errors=True))
|
||||
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
|
||||
held_by: Final = MappingProxyType({pid: _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)
|
||||
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)
|
||||
follow_up: Final = _Call("responses", False, uuid.uuid4().hex)
|
||||
(answered,) = await _burst(candidate, _CONFIG_MODEL, (follow_up,))
|
||||
_assert_answered_in_its_own_shape(answered)
|
||||
_assert_forwarded(_marked(wire.drain()), (*calls, follow_up))
|
||||
|
|
@ -0,0 +1,456 @@
|
|||
import json
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support import agentic_probe as ap
|
||||
from integration._support import codex_vendor as cv
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.wire import Wire, wire_server
|
||||
from openai.types.responses import (
|
||||
EasyInputMessageParam,
|
||||
ResponseCompletedEvent,
|
||||
ResponseInputItemParam,
|
||||
ResponseTextDeltaEvent,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(240)
|
||||
|
||||
_MODEL: Final = "chatgpt/gpt-5.5"
|
||||
_NO_CACHE: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}}
|
||||
_INCOMPLETE_GATE: Final = threading.Event()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Rig:
|
||||
wire: Wire
|
||||
proxy: OwnedProxy
|
||||
probe: Path
|
||||
|
||||
@property
|
||||
def gateway(self) -> Gateway:
|
||||
return self.proxy.gateway
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
return str(self.proxy.gateway.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
|
||||
directory: Final = tmp_path_factory.mktemp("chatgpt-caller-stream-rig")
|
||||
probe: Final = directory / "agentic-probe.jsonl"
|
||||
probe.touch()
|
||||
vendor: Final = cv.CodexVendor(incomplete_gate=_INCOMPLETE_GATE)
|
||||
with gateway_from_environment() as gateway, wire_server(vendor.respond) as wire:
|
||||
overrides: Final = {
|
||||
"CHATGPT_TOKEN_DIR": str(cv.login(directory)),
|
||||
"CHATGPT_API_BASE": wire.url,
|
||||
ap.OUT_ENVIRONMENT: str(probe),
|
||||
}
|
||||
config: Final = cv.proxy_config(directory, probe=True)
|
||||
with owned_proxy_process(gateway, directory, overrides, config=config, workers=2) as owned:
|
||||
yield _Rig(wire, owned, probe)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model(rig: _Rig) -> Iterator[str]:
|
||||
rig.wire.drain()
|
||||
with rig.gateway.scenario() as scenario:
|
||||
yield scenario.model(model=_MODEL, api_base=rig.wire.url, api_key=None)
|
||||
|
||||
|
||||
def _prompt(marker: str, *directives: str) -> str:
|
||||
return " ".join((f"Say marker-{marker}", *directives))
|
||||
|
||||
|
||||
def _sdk_input(marker: str) -> list[ResponseInputItemParam]: # mutable-ok: the OpenAI SDK input parameter is a list
|
||||
return [EasyInputMessageParam(role="user", content=_prompt(marker))]
|
||||
|
||||
|
||||
def _responses_body(
|
||||
model: str,
|
||||
marker: str,
|
||||
*directives: str,
|
||||
stream: bool | None = None,
|
||||
extra_body: Mapping[str, JsonValue] | None = None,
|
||||
cache_bust: bool = True,
|
||||
) -> Mapping[str, JsonValue]:
|
||||
return {
|
||||
"model": model,
|
||||
"input": [{"role": "user", "content": _prompt(marker, *directives)}],
|
||||
**(_NO_CACHE if cache_bust else {}),
|
||||
**({} if stream is None else {"stream": stream}),
|
||||
**({} if extra_body is None else {"extra_body": dict(extra_body)}),
|
||||
}
|
||||
|
||||
|
||||
def _post(rig: _Rig, path: str, body: Mapping[str, JsonValue]) -> httpx.Response:
|
||||
return rig.gateway.request("POST", path, body)
|
||||
|
||||
|
||||
def _only_forwarded(rig: _Rig, marker: str, *, stream: bool = True) -> Mapping[str, JsonValue]:
|
||||
(request,) = rig.wire.drain()
|
||||
return cv.forwarded(request, marker, stream=stream)
|
||||
|
||||
|
||||
def _spend_rows(model: str, expected: int) -> Sequence[Mapping[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, litellm_call_id, status, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE model_group = %s',
|
||||
(model,),
|
||||
),
|
||||
lambda rows: len(rows) >= expected,
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def _landed(model: str, response_id: str, *, expected: int = 1) -> Sequence[Mapping[str, JsonValue]]:
|
||||
rows: Final = _spend_rows(model, expected)
|
||||
(row,) = [row for row in rows if rv.same_response(string_value(row["request_id"]), response_id)]
|
||||
assert row["status"] == "success", rows
|
||||
return rows
|
||||
|
||||
|
||||
def _landed_by_call_id(model: str, call_id: str) -> None:
|
||||
rows: Final = _spend_rows(model, 1)
|
||||
(row,) = [row for row in rows if row["litellm_call_id"] == call_id]
|
||||
assert row["status"] == "success", rows
|
||||
|
||||
|
||||
def _output_text(body: Mapping[str, JsonValue]) -> str:
|
||||
(item,) = rv.ITEMS.validate_python(body["output"])
|
||||
(part,) = rv.ITEMS.validate_python(item["content"])
|
||||
return string_value(part["text"])
|
||||
|
||||
|
||||
def _assert_json_answer(response: httpx.Response, marker: str) -> str:
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.headers["content-type"].startswith("application/json"), response.headers
|
||||
assert "event:" not in response.text, response.text
|
||||
body: Final = rv.JSON_OBJECT.validate_json(response.content)
|
||||
assert _output_text(body) == rv.answer(marker), body
|
||||
assert cv.totals(rv.JSON_OBJECT.validate_python(body["usage"])) == cv.totals(cv.USAGE), body
|
||||
return string_value(body["id"])
|
||||
|
||||
|
||||
def _sse_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 _assert_sse_answer(response: httpx.Response, marker: str) -> str:
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.headers["content-type"].startswith("text/event-stream"), response.headers
|
||||
frames: Final = _sse_frames(response.text)
|
||||
(completed,) = [frame for frame in frames if frame.get("type") == "response.completed"]
|
||||
deltas: Final = "".join(
|
||||
string_value(frame["delta"]) for frame in frames if frame.get("type") == "response.output_text.delta"
|
||||
)
|
||||
assert deltas == rv.answer(marker), response.text
|
||||
return string_value(rv.JSON_OBJECT.validate_python(completed["response"])["id"])
|
||||
|
||||
|
||||
def _assert_error(response: httpx.Response, *, status: int | None, message: str) -> None:
|
||||
assert response.status_code >= 400, response.text
|
||||
assert status is None or response.status_code == status, response.text
|
||||
assert response.headers["content-type"].startswith("application/json"), response.headers
|
||||
assert "event:" not in response.text, response.text
|
||||
assert message in response.text, response.text
|
||||
|
||||
|
||||
def _openai(rig: _Rig) -> openai.OpenAI:
|
||||
return openai.OpenAI(base_url=f"{rig.base_url}/v1", api_key=rig.gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _async_openai(rig: _Rig) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(base_url=f"{rig.base_url}/v1", api_key=rig.gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _anthropic(rig: _Rig) -> anthropic.Anthropic:
|
||||
return anthropic.Anthropic(base_url=rig.base_url, api_key=rig.gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def test_openai_sdk_request_without_a_stream_flag_gets_the_aggregated_json_response(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
raw: Final = _openai(rig).responses.with_raw_response.create(
|
||||
model=model, input=_sdk_input(marker), extra_body=_NO_CACHE
|
||||
)
|
||||
assert raw.headers["content-type"].startswith("application/json"), raw.headers
|
||||
response: Final = raw.parse()
|
||||
assert response.output_text == rv.answer(marker), raw.text
|
||||
assert response.usage is not None and response.usage.total_tokens == cv.USAGE["total_tokens"], raw.text
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, response.id)
|
||||
|
||||
|
||||
async def test_async_openai_sdk_request_with_stream_false_gets_the_aggregated_json_response(
|
||||
rig: _Rig, model: str
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
raw: Final = await _async_openai(rig).responses.with_raw_response.create(
|
||||
model=model, input=_sdk_input(marker), stream=False, extra_body=_NO_CACHE
|
||||
)
|
||||
assert raw.headers["content-type"].startswith("application/json"), raw.headers
|
||||
response: Final = raw.parse()
|
||||
assert response.output_text == rv.answer(marker), raw.text
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, response.id)
|
||||
|
||||
|
||||
def test_raw_request_without_a_stream_key_gets_json_not_sse(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker))
|
||||
identity: Final = _assert_json_answer(response, marker)
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, identity)
|
||||
|
||||
|
||||
def test_openai_sdk_stream_request_still_streams(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
events: Final = list(
|
||||
_openai(rig).responses.create(model=model, input=_sdk_input(marker), stream=True, extra_body=_NO_CACHE)
|
||||
)
|
||||
deltas: Final = "".join(event.delta for event in events if isinstance(event, ResponseTextDeltaEvent))
|
||||
assert deltas == rv.answer(marker), events
|
||||
(completed,) = [event for event in events if isinstance(event, ResponseCompletedEvent)]
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, completed.response.id)
|
||||
|
||||
|
||||
async def test_async_openai_sdk_stream_request_still_streams(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
stream: Final = await _async_openai(rig).responses.create(
|
||||
model=model, input=_sdk_input(marker), stream=True, extra_body=_NO_CACHE
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
deltas: Final = "".join(event.delta for event in events if isinstance(event, ResponseTextDeltaEvent))
|
||||
assert deltas == rv.answer(marker), events
|
||||
(completed,) = [event for event in events if isinstance(event, ResponseCompletedEvent)]
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, completed.response.id)
|
||||
|
||||
|
||||
def test_openai_sdk_chat_completion_is_bridged_to_a_json_answer(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
completion: Final = _openai(rig).chat.completions.create(
|
||||
model=model, messages=[{"role": "user", "content": _prompt(marker)}], extra_body=_NO_CACHE
|
||||
)
|
||||
assert completion.choices[0].message.content == rv.answer(marker), completion
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, completion.id)
|
||||
|
||||
|
||||
async def test_async_openai_sdk_chat_completion_is_bridged_to_a_json_answer(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
completion: Final = await _async_openai(rig).chat.completions.create(
|
||||
model=model, messages=[{"role": "user", "content": _prompt(marker)}], extra_body=_NO_CACHE
|
||||
)
|
||||
assert completion.choices[0].message.content == rv.answer(marker), completion
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, completion.id)
|
||||
|
||||
|
||||
def test_openai_sdk_chat_completion_stream_still_streams(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
chunks: Final = list(
|
||||
_openai(rig).chat.completions.create(
|
||||
model=model, messages=[{"role": "user", "content": _prompt(marker)}], stream=True, extra_body=_NO_CACHE
|
||||
)
|
||||
)
|
||||
text: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
|
||||
assert text == rv.answer(marker), chunks
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, chunks[0].id)
|
||||
|
||||
|
||||
def test_anthropic_sdk_message_is_bridged_to_a_json_answer(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
message: Final = _anthropic(rig).messages.create(
|
||||
model=model, max_tokens=64, messages=[{"role": "user", "content": _prompt(marker)}], extra_body=_NO_CACHE
|
||||
)
|
||||
(block,) = message.content
|
||||
assert block.type == "text" and block.text == rv.answer(marker), message
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, message.id)
|
||||
|
||||
|
||||
def test_anthropic_sdk_message_stream_still_streams(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with _anthropic(rig).messages.stream(
|
||||
model=model, max_tokens=64, messages=[{"role": "user", "content": _prompt(marker)}], extra_body=_NO_CACHE
|
||||
) as stream:
|
||||
text: Final = "".join(stream.text_stream)
|
||||
message: Final = stream.get_final_message()
|
||||
assert text == rv.answer(marker), message
|
||||
_only_forwarded(rig, marker)
|
||||
_landed_by_call_id(model, stream.response.headers["x-litellm-call-id"])
|
||||
|
||||
|
||||
def test_identical_request_is_served_from_the_response_cache_as_json(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
body: Final = _responses_body(model, marker, cache_bust=False)
|
||||
first: Final = _post(rig, "/v1/responses", body)
|
||||
identity: Final = _assert_json_answer(first, marker)
|
||||
assert "x-litellm-cache-key" not in first.headers, first.headers
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, identity)
|
||||
second: Final = eventually(
|
||||
lambda: _post(rig, "/v1/responses", body), lambda found: "x-litellm-cache-key" in found.headers, seconds=20
|
||||
)
|
||||
assert rv.same_response(_assert_json_answer(second, marker), identity), second.text
|
||||
assert rig.wire.drain() == (), "the cached answer reached the vendor"
|
||||
rows: Final = _landed(model, identity, expected=2)
|
||||
(cached,) = [row for row in rows if "_cache_hit" in string_value(row["request_id"])]
|
||||
assert rv.same_response(string_value(cached["request_id"]).split("_cache_hit")[0], identity), rows
|
||||
assert (cached["status"], cached["cache_hit"]) == ("success", "True"), rows
|
||||
assert cached["spend"] == 0, rows
|
||||
|
||||
|
||||
def test_agentic_hook_sees_the_aggregated_response_once(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker))
|
||||
identity: Final = _assert_json_answer(response, marker)
|
||||
_landed(model, identity)
|
||||
(line,) = eventually(lambda: ap.lines(rig.probe, marker), lambda found: len(found) >= 1)
|
||||
assert (line["surface"], line["response_type"], line["stream"], line["provider"]) == (
|
||||
"responses",
|
||||
"ResponsesAPIResponse",
|
||||
False,
|
||||
"chatgpt",
|
||||
), line
|
||||
assert ap.lines(rig.probe, marker) == (line,), ap.lines(rig.probe, marker)
|
||||
|
||||
|
||||
def test_agentic_hook_stays_out_of_a_stream_request(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker, stream=True))
|
||||
identity: Final = _assert_sse_answer(response, marker)
|
||||
_landed(model, identity)
|
||||
assert ap.lines(rig.probe, marker) == (), ap.lines(rig.probe, marker)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Junk:
|
||||
label: str
|
||||
fragment: str
|
||||
streams: bool
|
||||
|
||||
|
||||
_JUNK: Final = (
|
||||
_Junk("null", '"stream": null', False),
|
||||
_Junk("empty-string", '"stream": ""', False),
|
||||
_Junk("empty-list", '"stream": []', False),
|
||||
_Junk("zero", '"stream": 0', False),
|
||||
_Junk("false-twice", '"stream": false, "stream": false', False),
|
||||
_Junk("true-then-false", '"stream": true, "stream": false', False),
|
||||
_Junk("one", '"stream": 1', True),
|
||||
_Junk("string-false", '"stream": "false"', True),
|
||||
_Junk("five-kb-string", f'"stream": "{"x" * 5000}"', True),
|
||||
_Junk("true-twice", '"stream": true, "stream": true', True),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("junk", _JUNK, ids=[junk.label for junk in _JUNK])
|
||||
def test_odd_stream_values_decide_the_shape_by_their_truth(rig: _Rig, model: str, junk: _Junk) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
body: Final = json.dumps(_responses_body(model, marker))[:-1] + f", {junk.fragment}}}"
|
||||
response: Final = rig.gateway.client.post(
|
||||
"/v1/responses",
|
||||
content=body.encode(),
|
||||
headers={"Authorization": f"Bearer {rig.gateway.key}", "content-type": "application/json"},
|
||||
)
|
||||
identity: Final = _assert_sse_answer(response, marker) if junk.streams else _assert_json_answer(response, marker)
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, identity)
|
||||
|
||||
|
||||
def test_extra_body_stream_true_streams(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker, extra_body={"stream": True}))
|
||||
identity: Final = _assert_sse_answer(response, marker)
|
||||
_only_forwarded(rig, marker)
|
||||
_landed(model, identity)
|
||||
|
||||
|
||||
def test_extra_body_stream_false_gets_json(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker, extra_body={"stream": False}))
|
||||
identity: Final = _assert_json_answer(response, marker)
|
||||
_only_forwarded(rig, marker, stream=False)
|
||||
_landed(model, identity)
|
||||
|
||||
|
||||
def test_string_input_is_refused_by_the_vendor_as_a_400(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", {"model": model, "input": _prompt(marker), **_NO_CACHE})
|
||||
_assert_error(response, status=400, message=cv.INPUT_MUST_BE_A_LIST)
|
||||
(request,) = rig.wire.drain()
|
||||
assert rv.JSON_OBJECT.validate_json(request.body)["input"] == _prompt(marker), request.body
|
||||
|
||||
|
||||
def test_vendor_401_reaches_the_caller_as_401(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker, cv.UNAUTHORIZED_DIRECTIVE))
|
||||
_assert_error(response, status=401, message=cv.UNAUTHORIZED)
|
||||
_only_forwarded(rig, marker)
|
||||
|
||||
|
||||
def test_vendor_response_failed_event_reaches_the_caller_as_a_json_error(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
response: Final = _post(rig, "/v1/responses", _responses_body(model, marker, cv.FAILED_DIRECTIVE))
|
||||
_assert_error(response, status=None, message=cv.failure_message(marker))
|
||||
_only_forwarded(rig, marker)
|
||||
|
||||
|
||||
def test_vendor_stream_dying_mid_transfer_answers_a_json_error_and_leaves_the_proxy_serving(
|
||||
rig: _Rig, model: str
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
_INCOMPLETE_GATE.clear()
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
pending: Final = pool.submit(
|
||||
_post, rig, "/v1/responses", _responses_body(model, marker, cv.INCOMPLETE_DIRECTIVE)
|
||||
)
|
||||
eventually(rig.wire.received.qsize, lambda size: size >= 1)
|
||||
liveliness: Final = rig.gateway.request("GET", "/health/liveliness", None)
|
||||
assert liveliness.status_code == 200, liveliness.text
|
||||
assert not pending.done(), pending.result().text
|
||||
_INCOMPLETE_GATE.set()
|
||||
response: Final = pending.result(timeout=60)
|
||||
assert response.status_code >= 400, response.text
|
||||
assert response.headers["content-type"].startswith("application/json"), response.headers
|
||||
assert "event:" not in response.text, response.text
|
||||
assert "error" in rv.JSON_OBJECT.validate_json(response.content), response.text
|
||||
_only_forwarded(rig, marker)
|
||||
follow_up: Final = uuid.uuid4().hex
|
||||
identity: Final = _assert_json_answer(_post(rig, "/v1/responses", _responses_body(model, follow_up)), follow_up)
|
||||
_only_forwarded(rig, follow_up)
|
||||
_landed(model, identity, expected=2)
|
||||
|
||||
|
||||
def test_repeated_identical_cache_busted_requests_each_land_once(rig: _Rig, model: str) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
body: Final = _responses_body(model, marker)
|
||||
identities: Final = tuple(_assert_json_answer(_post(rig, "/v1/responses", body), marker) for _ in range(5))
|
||||
assert len(set(identities)) == 5, identities
|
||||
received: Final = rig.wire.drain()
|
||||
assert len(received) == 5, [request.target for request in received]
|
||||
for request in received:
|
||||
cv.forwarded(request, marker)
|
||||
rows: Final = _spend_rows(model, 5)
|
||||
assert len(rows) == 5, rows
|
||||
for identity in identities:
|
||||
(row,) = [row for row in rows if rv.same_response(string_value(row["request_id"]), identity)]
|
||||
assert row["status"] == "success", rows
|
||||
|
|
@ -31,6 +31,8 @@ from litellm.llms.brave.search.transformation import BraveSearchConfig
|
|||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
|
||||
from litellm.llms.chatgpt.authenticator import Authenticator as ChatGPTAuthenticator
|
||||
from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPIConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
|
|
@ -50,6 +52,10 @@ from litellm.llms.openai.vector_store_files.transformation import OpenAIVectorSt
|
|||
from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig
|
||||
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
|
||||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
from litellm.responses.streaming_iterator import (
|
||||
BaseResponsesAPIStreamingIterator,
|
||||
MockResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse
|
||||
|
|
@ -364,6 +370,7 @@ async def test_async_response_api_handler_streams_when_provider_transform_adds_s
|
|||
)
|
||||
)
|
||||
logging_obj = Mock()
|
||||
logging_obj.dynamic_success_callbacks = None
|
||||
|
||||
await handler.async_response_api_handler(
|
||||
model="gpt-5.3-codex",
|
||||
|
|
@ -380,6 +387,166 @@ async def test_async_response_api_handler_streams_when_provider_transform_adds_s
|
|||
assert client.post.call_args.kwargs["json"]["stream"] is True
|
||||
|
||||
|
||||
_CHATGPT_SSE_BODY = (
|
||||
"event: response.output_item.done\n"
|
||||
'data: {"type": "response.output_item.done", "output_index": 0, "item": {"type": "message", '
|
||||
'"id": "msg_1", "status": "completed", "role": "assistant", "content": [{"type": "output_text", '
|
||||
'"text": "aggregated", "annotations": []}]}}\n'
|
||||
"\n"
|
||||
"event: response.completed\n"
|
||||
'data: {"type": "response.completed", "response": {"id": "resp_1", "object": "response", '
|
||||
'"created_at": 1, "model": "gpt-5.3-codex", "status": "completed", "output": [], '
|
||||
'"parallel_tool_calls": false, "tool_choice": "auto", "tools": [], '
|
||||
'"usage": {"input_tokens": 1, "output_tokens": 2, "total_tokens": 3}}}\n'
|
||||
"\n"
|
||||
)
|
||||
|
||||
|
||||
def _chatgpt_sse_response():
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
content=_CHATGPT_SSE_BODY.encode(),
|
||||
request=httpx.Request("POST", "https://chatgpt.example.com/responses"),
|
||||
)
|
||||
|
||||
|
||||
def _chatgpt_responses_logging_obj():
|
||||
logging_obj = Mock()
|
||||
logging_obj.dynamic_success_callbacks = None
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _chatgpt_responses_config():
|
||||
authenticator = Mock(spec=ChatGPTAuthenticator)
|
||||
authenticator.get_access_token.return_value = "access-test"
|
||||
authenticator.get_account_id.return_value = "acct-test"
|
||||
return ChatGPTResponsesAPIConfig(authenticator=authenticator)
|
||||
|
||||
|
||||
def _chatgpt_handler_kwargs(caller_params, client):
|
||||
return {
|
||||
"model": "gpt-5.3-codex",
|
||||
"input": "hi",
|
||||
"responses_api_provider_config": _chatgpt_responses_config(),
|
||||
"response_api_optional_request_params": caller_params,
|
||||
"custom_llm_provider": "chatgpt",
|
||||
"litellm_params": GenericLiteLLMParams(api_key="sk-test", api_base="https://chatgpt.example.com"),
|
||||
"logging_obj": _chatgpt_responses_logging_obj(),
|
||||
"client": client,
|
||||
}
|
||||
|
||||
|
||||
def _assert_aggregated_chatgpt_response(result):
|
||||
assert not isinstance(result, BaseResponsesAPIStreamingIterator)
|
||||
assert isinstance(result, ResponsesAPIResponse)
|
||||
assert result.id == "resp_1"
|
||||
assert result.output[0].content[0].text == "aggregated"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("caller_params", [{}, {"stream": False}])
|
||||
def test_response_api_handler_aggregates_chatgpt_sse_for_a_non_streaming_caller(caller_params):
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = HTTPHandler(client=httpx.Client())
|
||||
client.post = Mock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = handler.response_api_handler(**_chatgpt_handler_kwargs(caller_params, client))
|
||||
|
||||
assert client.post.call_args.kwargs["json"]["stream"] is True
|
||||
_assert_aggregated_chatgpt_response(result)
|
||||
|
||||
|
||||
def test_response_api_handler_streams_chatgpt_sse_for_a_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = HTTPHandler(client=httpx.Client())
|
||||
client.post = Mock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = handler.response_api_handler(**_chatgpt_handler_kwargs({"stream": True}, client))
|
||||
|
||||
assert isinstance(result, BaseResponsesAPIStreamingIterator)
|
||||
assert [event.type for event in result][-1] == "response.completed"
|
||||
|
||||
|
||||
def test_response_api_handler_streams_chatgpt_sse_for_an_extra_body_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = HTTPHandler(client=httpx.Client())
|
||||
client.post = Mock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = handler.response_api_handler(extra_body={"stream": True}, **_chatgpt_handler_kwargs({}, client))
|
||||
|
||||
assert isinstance(result, BaseResponsesAPIStreamingIterator)
|
||||
assert [event.type for event in result][-1] == "response.completed"
|
||||
|
||||
|
||||
def test_response_api_handler_fake_streams_only_for_a_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = HTTPHandler(client=httpx.Client())
|
||||
client.post = Mock(return_value=_chatgpt_sse_response())
|
||||
|
||||
streamed = handler.response_api_handler(fake_stream=True, **_chatgpt_handler_kwargs({"stream": True}, client))
|
||||
assert isinstance(streamed, MockResponsesAPIStreamingIterator)
|
||||
|
||||
client.post = Mock(return_value=_chatgpt_sse_response())
|
||||
aggregated = handler.response_api_handler(fake_stream=True, **_chatgpt_handler_kwargs({}, client))
|
||||
_assert_aggregated_chatgpt_response(aggregated)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("caller_params", [{}, {"stream": False}])
|
||||
async def test_async_response_api_handler_aggregates_chatgpt_sse_for_a_non_streaming_caller(caller_params):
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = AsyncHTTPHandler()
|
||||
client.post = AsyncMock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = await handler.async_response_api_handler(**_chatgpt_handler_kwargs(caller_params, client))
|
||||
|
||||
assert client.post.call_args.kwargs["json"]["stream"] is True
|
||||
_assert_aggregated_chatgpt_response(result)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_response_api_handler_streams_chatgpt_sse_for_a_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = AsyncHTTPHandler()
|
||||
client.post = AsyncMock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = await handler.async_response_api_handler(**_chatgpt_handler_kwargs({"stream": True}, client))
|
||||
|
||||
assert isinstance(result, BaseResponsesAPIStreamingIterator)
|
||||
assert [event.type async for event in result][-1] == "response.completed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_response_api_handler_streams_chatgpt_sse_for_an_extra_body_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = AsyncHTTPHandler()
|
||||
client.post = AsyncMock(return_value=_chatgpt_sse_response())
|
||||
|
||||
result = await handler.async_response_api_handler(
|
||||
extra_body={"stream": True}, **_chatgpt_handler_kwargs({}, client)
|
||||
)
|
||||
|
||||
assert isinstance(result, BaseResponsesAPIStreamingIterator)
|
||||
assert [event.type async for event in result][-1] == "response.completed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_response_api_handler_fake_streams_only_for_a_streaming_caller():
|
||||
handler = BaseLLMHTTPHandler()
|
||||
client = AsyncHTTPHandler()
|
||||
client.post = AsyncMock(return_value=_chatgpt_sse_response())
|
||||
|
||||
streamed = await handler.async_response_api_handler(
|
||||
fake_stream=True, **_chatgpt_handler_kwargs({"stream": True}, client)
|
||||
)
|
||||
assert isinstance(streamed, MockResponsesAPIStreamingIterator)
|
||||
|
||||
client.post = AsyncMock(return_value=_chatgpt_sse_response())
|
||||
aggregated = await handler.async_response_api_handler(fake_stream=True, **_chatgpt_handler_kwargs({}, client))
|
||||
_assert_aggregated_chatgpt_response(aggregated)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_response_api_handler_streaming_passes_logging_obj_to_post():
|
||||
"""LIT-5466: @track_llm_api_timing only records llm_api_duration_ms when the POST
|
||||
|
|
@ -407,7 +574,7 @@ async def test_async_response_api_handler_streaming_passes_logging_obj_to_post()
|
|||
model="gpt-5",
|
||||
input="hi",
|
||||
responses_api_provider_config=config,
|
||||
response_api_optional_request_params={},
|
||||
response_api_optional_request_params={"stream": True},
|
||||
custom_llm_provider="chatgpt",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -441,7 +608,7 @@ async def test_async_response_api_handler_posts_the_async_transform_hook_result(
|
|||
model="gpt-5",
|
||||
input="hi",
|
||||
responses_api_provider_config=config,
|
||||
response_api_optional_request_params={},
|
||||
response_api_optional_request_params={"stream": True},
|
||||
custom_llm_provider="chatgpt",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=Mock(),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue