diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index 78ee8a86a65..d66e3d97d66 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -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: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7dd3b9b34c0..86c36a7b413 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/tests/integration/_support/agentic_probe.py b/tests/integration/_support/agentic_probe.py new file mode 100644 index 00000000000..ff264718d7d --- /dev/null +++ b/tests/integration/_support/agentic_probe.py @@ -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) diff --git a/tests/integration/_support/codex_vendor.py b/tests/integration/_support/codex_vendor.py new file mode 100644 index 00000000000..2ec9c4e8b96 --- /dev/null +++ b/tests/integration/_support/codex_vendor.py @@ -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) diff --git a/tests/integration/providers/test_chatgpt_responses_caller_stream_chaos.py b/tests/integration/providers/test_chatgpt_responses_caller_stream_chaos.py new file mode 100644 index 00000000000..d1e199c949a --- /dev/null +++ b/tests/integration/providers/test_chatgpt_responses_caller_stream_chaos.py @@ -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)) diff --git a/tests/integration/providers/test_chatgpt_responses_caller_stream_wire.py b/tests/integration/providers/test_chatgpt_responses_caller_stream_wire.py new file mode 100644 index 00000000000..d4ae7e0b4ae --- /dev/null +++ b/tests/integration/providers/test_chatgpt_responses_caller_stream_wire.py @@ -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 diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index 21c62500df3..93d18727f22 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -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(),