diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 1c34eda3..dd3580fe 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -94,3 +94,18 @@ jobs: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7 - run: bash scripts/release-tests/litellm-existing-secret.sh + + litellm-stream-errors: + name: LiteLLM Responses stream errors + runs-on: ubuntu-latest + timeout-minutes: 5 + steps: + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7 + + - uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7 + with: + python-version: "3.11" + + - run: python3 -m unittest discover -s deploy/litellm/tests -p 'test_*.py' + + - run: bash scripts/release-tests/litellm-build-context.sh diff --git a/deploy/litellm/Dockerfile b/deploy/litellm/Dockerfile index 4c058738..a40ac6b6 100644 --- a/deploy/litellm/Dockerfile +++ b/deploy/litellm/Dockerfile @@ -6,6 +6,10 @@ FROM ghcr.io/berriai/litellm:v1.99.0@sha256:570a872d2fde8f1bc4a147634810941103c14697270f0f2918c5aa7d8201cac5 WORKDIR /app + +COPY patches/ /tmp/litellm-patches/ +RUN python3 /tmp/litellm-patches/apply_responses_stream_errors.py && \ + rm -rf /tmp/litellm-patches # ═══ Inject custom hook files ═══ # In the base image, litellm is pip-installed, so the import path resolves to @@ -36,4 +40,4 @@ RUN chmod +x ./docker/entrypoint.sh EXPOSE 4000/tcp -CMD ["--port", "4000", "--config", "config.yaml"] \ No newline at end of file +CMD ["--port", "4000", "--config", "config.yaml"] diff --git a/deploy/litellm/README.md b/deploy/litellm/README.md index 46a962ff..3ae3652c 100644 --- a/deploy/litellm/README.md +++ b/deploy/litellm/README.md @@ -114,11 +114,53 @@ Secret in that namespace and set `serviceMonitor.bearerTokenSecret.name`. ## Building the image -`./build.sh` (needs `REGISTRY`; set `HARBOR_PASSWORD` to build in-cluster with -kaniko when there is no docker daemon). The tag it writes names the LiteLLM +`./build.sh` needs `REGISTRY`. Set `PUSH_SECRET` to select an in-cluster Kaniko +build using an existing registry credential Secret when there is no Docker +daemon. The tag it writes names the LiteLLM version taken from the Dockerfile, and `deploy.sh` refuses an image whose version-named tag disagrees with that pin. +## Responses stream errors + +The image includes a checked patch for LiteLLM 1.99.0. When an upstream request +fails after a native Responses stream has started, the proxy emits +`response.failed` with the upstream error message and a recognizable error code. +This lets Responses clients report the cause instead of an unexpected end of +stream. The patch preserves the response ID and event sequence when available, +and does not append another failure after a terminal event. Chat Completions +and the Cursor conversion endpoint keep their existing stream formats. + +The build verifies SHA-256 hashes of the upstream files before applying the +patch. A different source version fails the build and requires reviewing the +patch against that version. Both Docker and Kaniko include the patch installer +and helper in their build context. Retry, fallback, and cooldown policies are +unchanged by this patch. + +From the repository root, run the local checks: + +```bash +python3 -m unittest discover -s deploy/litellm/tests -p 'test_*.py' +bash scripts/release-tests/litellm-build-context.sh +``` + +Verify a built image against a synthetic upstream without provider credentials +or a database: + +```bash +docker run --rm --entrypoint python3 \ + -v "$PWD/deploy/litellm/tests:/tests:ro" "$LITELLM_TEST_IMAGE" \ + /tests/integration_responses_stream_errors.py +``` + +The integration check exercises the actual LiteLLM HTTP server. Its `--serve` +mode keeps the synthetic service running for a Codex CLI check through a local +port-forward: + +```bash +python3 deploy/litellm/tests/integration_responses_stream_errors.py \ + --codex-url http://127.0.0.1:4000 +``` + ## Upgrading LiteLLM **The database migration is one-way.** LiteLLM runs `prisma migrate deploy` on diff --git a/deploy/litellm/build.sh b/deploy/litellm/build.sh index cc3905ed..b96bac03 100755 --- a/deploy/litellm/build.sh +++ b/deploy/litellm/build.sh @@ -10,9 +10,8 @@ # away from the Dockerfile beside it. `claw/deploy/build.sh` had solved the # same problem with an in-cluster kaniko job; this borrows that shape. # -# HARBOR_PASSWORD push password for $REGISTRY (presence selects kaniko) -# PUSH_SECRET Secret holding .dockerconfigjson for $REGISTRY (kaniko backend) -# HARBOR_USERNAME push user (default: admin) +# PUSH_SECRET Secret holding .dockerconfigjson; selects the kaniko backend +# HARBOR_PASSWORD legacy selector for kaniko; credentials still use PUSH_SECRET # REGISTRY e.g. harbor.example.com/primussafe # NAMESPACE where the build job runs (default: primus-claw) # TAG default: v- @@ -26,7 +25,6 @@ SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" NAMESPACE="${NAMESPACE:-primus-claw}" REGISTRY="${REGISTRY:?REGISTRY is required, e.g. harbor.example.com/primussafe}" -HARBOR_USERNAME="${HARBOR_USERNAME:-admin}" # Read the pinned version out of the Dockerfile so the tag cannot disagree. BASE_VERSION="$(grep -oE '^FROM .*litellm:v[0-9.]+' "$SCRIPT_DIR/Dockerfile" | grep -oE 'v[0-9.]+' | head -1)" @@ -36,8 +34,8 @@ IMG="$REGISTRY/litellm:$TAG" echo "[litellm-build] building $IMG (base $BASE_VERSION)" -if [ -z "${HARBOR_PASSWORD:-}" ]; then - command -v docker >/dev/null || { echo "ERROR: no HARBOR_PASSWORD for the kaniko backend and no docker daemon" >&2; exit 1; } +if [ -z "${PUSH_SECRET:-}" ] && [ -z "${HARBOR_PASSWORD:-}" ]; then + command -v docker >/dev/null || { echo "ERROR: set PUSH_SECRET for kaniko, or install docker" >&2; exit 1; } docker build -t "$IMG" "$SCRIPT_DIR" docker push "$IMG" echo "[litellm-build] pushed $IMG" @@ -46,12 +44,17 @@ fi JOB="litellm-build-$(date +%s)" CTX="litellm-build-ctx-$(date +%s)" -cleanup() { kubectl -n "$NAMESPACE" delete cm "$CTX" --ignore-not-found >/dev/null 2>&1 || true; } +CONTEXT_ARCHIVE="$(mktemp)" +cleanup() { + rm -f "$CONTEXT_ARCHIVE" + kubectl -n "$NAMESPACE" delete cm "$CTX" --ignore-not-found >/dev/null 2>&1 || true +} trap cleanup EXIT +tar -C "$SCRIPT_DIR" --exclude='__pycache__' --exclude='*.pyc' \ + -czf "$CONTEXT_ARCHIVE" Dockerfile apim_key_hook.py patches kubectl -n "$NAMESPACE" create cm "$CTX" \ - --from-file=Dockerfile="$SCRIPT_DIR/Dockerfile" \ - --from-file=apim_key_hook.py="$SCRIPT_DIR/apim_key_hook.py" \ + --from-file=build-context.tar.gz="$CONTEXT_ARCHIVE" \ --dry-run=client -o yaml | kubectl apply -f - >/dev/null # Registry auth is read from an existing pull/push secret rather than inlined @@ -75,7 +78,7 @@ spec: initContainers: - name: ctx image: busybox:1.36 - command: ["sh","-c","cp /cm/Dockerfile /cm/apim_key_hook.py /workspace/"] + command: ["sh","-c","tar -xzf /cm/build-context.tar.gz -C /workspace"] volumeMounts: - {name: ws, mountPath: /workspace} - {name: cm, mountPath: /cm} diff --git a/deploy/litellm/patches/apply_responses_stream_errors.py b/deploy/litellm/patches/apply_responses_stream_errors.py new file mode 100644 index 00000000..0608d72b --- /dev/null +++ b/deploy/litellm/patches/apply_responses_stream_errors.py @@ -0,0 +1,173 @@ +# Copyright Advanced Micro Devices, Inc. +# SPDX-License-Identifier: MIT + +from __future__ import annotations + +import argparse +import ast +import hashlib +import importlib.util +from pathlib import Path + + +BASE_SHA256 = { + "proxy/proxy_server.py": "f63cd83c5c4459d84dbfbd350a460caf9e2907389f19d2c3090fe24ad937f47b", + "proxy/response_api_endpoints/endpoints.py": "563b462e7c36e729869d0acdc016a54c2e99ec07f70a3f6ea17482b1dbf11e4f", + "responses/streaming_iterator.py": "e25f32c3bf3e1815f08a5f2e328d4b32359d189d386780a7b9b4640bdbe4e56a", +} +HELPER_NAME = "responses_stream_errors.py" + + +def _replace_once(source: str, before: str, after: str) -> str: + count = source.count(before) + if count != 1: + raise ValueError(f"Expected exactly one patch context, found {count}: {before[:80]!r}") + return source.replace(before, after, 1) + + +def _native_responses_endpoint(source: str) -> str: + function = next(node for node in ast.parse(source).body if isinstance(node, ast.AsyncFunctionDef) and node.name == "responses_api") + lines = source.splitlines(keepends=True) + start, end = function.lineno - 1, function.end_lineno + block = "".join(lines[start:end]) + block = _replace_once(block, " select_data_generator,\n", " select_responses_data_generator as select_data_generator,\n") + return "".join(lines[:start]) + block + "".join(lines[end:]) + + +def _responses_iterator(source: str) -> str: + before = """ if 400 <= status_code < 500 and status_code != 429: + raise mapped_exception +""" + after = """ from litellm.proxy.responses_stream_errors import preserve_upstream_error + + preserve_upstream_error(mapped_exception, error_obj, result) + if 400 <= status_code < 500 and status_code != 429: + raise mapped_exception +""" + return _replace_once(source, before, after) + + +def _proxy_generator(source: str) -> str: + before = """async def async_data_generator( + response, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + request: Request | None = None, +): + verbose_proxy_logger.debug("inside generator") + stream_completed = False + client_disconnected = False +""" + after = """async def async_data_generator( + response, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + request: Request | None = None, + *, + responses_stream_errors: bool = False, +): + from litellm.proxy.responses_stream_errors import ResponsesStreamErrorState + + verbose_proxy_logger.debug("inside generator") + stream_completed = False + client_disconnected = False + error_state = ResponsesStreamErrorState() if responses_stream_errors else None +""" + source = _replace_once(source, before, after) + source = _replace_once(source, " raw_passthrough = False\n", " if error_state is not None:\n error_state.observe_chunk(chunk, emitted=False)\n raw_passthrough = False\n") + before = """ if isinstance(e, HTTPException): + raise e + elif isinstance(e, StreamingCallbackError): +""" + after = """ if error_state is not None: + stream_completed = True + error_frame = error_state.format_failure(e) + if error_frame is not None: + yield error_frame + return + if isinstance(e, HTTPException): + raise e + elif isinstance(e, StreamingCallbackError): +""" + source = _replace_once(source, before, after) + before = ''' yield _format_streaming_sse_chunk(chunk=chunk) + except Exception as e: + yield f"data: {e}\\n\\n" +''' + after = ''' formatted_chunk = _format_streaming_sse_chunk(chunk=chunk) + if error_state is not None: + error_state.mark_emitted() + yield formatted_chunk + except Exception as e: + if error_state is not None: + raise + yield f"data: {e}\\n\\n" +''' + return _replace_once(source, before, after) + + +def _responses_selector(source: str) -> str: + before = "\ndef select_data_generator(\n" + after = """ +def select_responses_data_generator( + response, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + request: Request | None = None, +): + return async_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + request=request, + responses_stream_errors=True, + ) + + +def select_data_generator( +""" + return _replace_once(source, before, after) + + +def apply_patch(package_root: Path) -> None: + sources: dict[str, str] = {} + for relative, expected in BASE_SHA256.items(): + content = (package_root / relative).read_bytes() + actual = hashlib.sha256(content).hexdigest() + if actual != expected: + raise ValueError(f"Unsupported LiteLLM source {relative}: expected SHA-256 {expected}, found {actual}") + sources[relative] = content.decode("utf-8") + helper = Path(__file__).with_name(HELPER_NAME).read_text() + target_helper = package_root / "proxy" / HELPER_NAME + if target_helper.exists(): + raise ValueError(f"Refusing to replace an existing {target_helper}") + sources["proxy/proxy_server.py"] = _responses_selector(_proxy_generator(sources["proxy/proxy_server.py"])) + sources["proxy/response_api_endpoints/endpoints.py"] = _native_responses_endpoint(sources["proxy/response_api_endpoints/endpoints.py"]) + sources["responses/streaming_iterator.py"] = _responses_iterator(sources["responses/streaming_iterator.py"]) + for relative, source in sources.items(): + compile(source, relative, "exec") + compile(helper, HELPER_NAME, "exec") + for relative, source in sources.items(): + (package_root / relative).write_text(source) + target_helper.write_text(helper) + + +def main() -> None: + parser = argparse.ArgumentParser(description="Apply the checked Responses stream error patch to LiteLLM 1.99.0.") + parser.add_argument("package_root", nargs="?", type=Path) + parser.add_argument("--package-dir", type=Path) + args = parser.parse_args() + if args.package_root is not None and args.package_dir is not None: + parser.error("Specify either package_root or --package-dir, not both") + package_root = args.package_dir or args.package_root + if package_root is None: + spec = importlib.util.find_spec("litellm") + if spec is None or spec.origin is None: + raise RuntimeError("Cannot locate the installed litellm package") + package_root = Path(spec.origin).parent + apply_patch(package_root) + print("Applied the Responses stream error patch to LiteLLM 1.99.0.") + + +if __name__ == "__main__": + main() diff --git a/deploy/litellm/patches/responses_stream_errors.py b/deploy/litellm/patches/responses_stream_errors.py new file mode 100644 index 00000000..12214a7f --- /dev/null +++ b/deploy/litellm/patches/responses_stream_errors.py @@ -0,0 +1,127 @@ +# Copyright Advanced Micro Devices, Inc. +# SPDX-License-Identifier: MIT + +from __future__ import annotations + +import json +import time +import uuid +from collections.abc import Mapping +from typing import Any + + +def _field(value: object, name: str, default: Any = None) -> Any: + if isinstance(value, Mapping): + return value.get(name, default) + return getattr(value, name, default) + + +def _response_fields(response: object) -> dict[str, Any]: + fields: dict[str, Any] = {} + for name, expected_type in (("id", str), ("model", str), ("created_at", int)): + value = _field(response, name) + if type(value) is expected_type and value != "": + fields[name] = value + return fields + + +def preserve_upstream_error(exception: Exception, error: object, event: object = None) -> None: + fields = { + name: value + for name in ("message", "code", "type", "param") + if isinstance(value := _field(error, name), (str, int)) + and not isinstance(value, bool) + } + fields["_response_fields"] = _response_fields(_field(event, "response")) + sequence_number = _field(event, "sequence_number") + if type(sequence_number) is int and sequence_number >= 0: + fields["_sequence_number"] = sequence_number + setattr(exception, "_responses_stream_error", fields) + + +def _upstream_error(exception: Exception) -> dict[str, Any]: + seen: set[int] = set() + current: object = exception + while isinstance(current, Exception) and id(current) not in seen: + seen.add(id(current)) + error = getattr(current, "_responses_stream_error", None) + if isinstance(error, dict) and error: + return error + current = getattr(current, "original_exception", None) + return {} + + +def _error_code(exception: Exception, upstream: Mapping[str, Any]) -> str: + code = upstream.get("code") + error_type = upstream.get("type") + for value in (code, error_type): + if isinstance(value, (str, int)): + normalized = str(value).lower() + if normalized == "insufficient_quota": + return "insufficient_quota" + if normalized in ("429", "toomanyrequests", "too_many_requests") or normalized.startswith("rate_limit"): + return "rate_limit_exceeded" + if isinstance(code, str) and code and not code.isdecimal(): + return code + status = getattr(exception, "status_code", None) + if str(status) == "429": + return "rate_limit_exceeded" + return { + "400": "invalid_request_error", + "401": "authentication_error", + "403": "permission_denied", + "404": "not_found_error", + "408": "request_timeout", + "422": "invalid_request_error", + }.get(str(status), "server_error") + + +class ResponsesStreamErrorState: + def __init__(self) -> None: + self.response_fields: dict[str, Any] = {} + self.sequence_number = -1 + self.terminal_seen = False + self._pending_terminal = False + + def observe_chunk(self, chunk: object, *, emitted: bool = True) -> None: + sequence_number = _field(chunk, "sequence_number") + if type(sequence_number) is int and sequence_number >= 0: + self.sequence_number = max(self.sequence_number, sequence_number) + self.response_fields.update(_response_fields(_field(chunk, "response"))) + response_id = _field(chunk, "response_id") + if not self.response_fields.get("id") and isinstance(response_id, str) and response_id: + self.response_fields["id"] = response_id + self._pending_terminal = _field(chunk, "type") in ("response.completed", "response.failed", "response.incomplete") + if emitted: + self.mark_emitted() + + def mark_emitted(self) -> None: + self.terminal_seen = self.terminal_seen or self._pending_terminal + self._pending_terminal = False + + def format_failure(self, exception: Exception) -> str | None: + if self.terminal_seen: + return None + upstream = _upstream_error(exception) + for name, value in upstream.get("_response_fields", {}).items(): + self.response_fields.setdefault(name, value) + upstream_sequence = upstream.get("_sequence_number") + if type(upstream_sequence) is int: + self.sequence_number = max(self.sequence_number, upstream_sequence - 1) + message = upstream.get("message") or getattr(exception, "message", None) or str(exception) + response = { + "id": "resp_" + uuid.uuid4().hex, + "object": "response", + "created_at": int(time.time()), + **self.response_fields, + "status": "failed", + "output": [], + "error": { + "code": _error_code(exception, upstream), + "message": str(message) or "The response could not be completed.", + }, + } + self.sequence_number += 1 + self.terminal_seen = True + payload = {"type": "response.failed", "sequence_number": self.sequence_number, "response": response} + return "event: response.failed\ndata: " + json.dumps(payload, separators=(",", ":")) + "\n\n" diff --git a/deploy/litellm/tests/integration_responses_stream_errors.py b/deploy/litellm/tests/integration_responses_stream_errors.py new file mode 100644 index 00000000..4946b4dc --- /dev/null +++ b/deploy/litellm/tests/integration_responses_stream_errors.py @@ -0,0 +1,420 @@ +#!/usr/bin/env python3 +# Copyright Advanced Micro Devices, Inc. +# SPDX-License-Identifier: MIT + +"""Exercise the installed LiteLLM proxy against a local synthetic upstream.""" + +import argparse +import contextlib +import http.server +import json +import os +from pathlib import Path +import re +import shutil +import socket +import subprocess +import tempfile +import threading +import time +import urllib.error +import urllib.parse +import urllib.request + + +MODEL = "gpt-5.2" +SUCCESS = "Synthetic completion." +ERROR_MESSAGE = "Synthetic upstream capacity error." +TOOL_ARGUMENTS = '{"value":"ok"}' + + +def response(status="in_progress", output=None): + return { + "id": "resp_fixture", + "object": "response", + "created_at": 1, + "status": status, + "error": None, + "incomplete_details": None, + "instructions": None, + "model": MODEL, + "output": output or [], + "parallel_tool_calls": True, + "tools": [], + "tool_choice": "auto", + "temperature": 1.0, + "top_p": 1.0, + "metadata": {}, + "usage": {"input_tokens": 1, "output_tokens": 2, "total_tokens": 3}, + } + + +def event(event_type, sequence, **fields): + return {"type": event_type, "sequence_number": sequence, **fields} + + +def tool_events(): + item = { + "id": "fc_fixture", + "type": "function_call", + "call_id": "call_fixture", + "name": "fixture_tool", + "arguments": TOOL_ARGUMENTS, + "status": "completed", + } + return [ + event("response.created", 0, response=response()), + event("response.output_item.added", 1, output_index=0, + item={**item, "arguments": "", "status": "in_progress"}), + event("response.function_call_arguments.delta", 2, + item_id=item["id"], output_index=0, delta=TOOL_ARGUMENTS), + event("response.function_call_arguments.done", 3, + item_id=item["id"], output_index=0, arguments=TOOL_ARGUMENTS), + event("response.output_item.done", 4, output_index=0, item=item), + event("response.completed", 5, response=response("completed", [item])), + ] + + +def message_events(): + part = {"type": "output_text", "text": SUCCESS, "annotations": []} + item = {"id": "msg_fixture", "type": "message", "role": "assistant", + "status": "completed", "content": [part]} + return [ + event("response.created", 0, response=response()), + event("response.output_item.added", 1, output_index=0, + item={**item, "status": "in_progress", "content": []}), + event("response.content_part.added", 2, item_id=item["id"], + output_index=0, content_index=0, part={**part, "text": ""}), + event("response.output_text.delta", 3, item_id=item["id"], + output_index=0, content_index=0, delta=SUCCESS, logprobs=[]), + event("response.output_text.done", 4, item_id=item["id"], + output_index=0, content_index=0, text=SUCCESS, logprobs=[]), + event("response.content_part.done", 5, item_id=item["id"], + output_index=0, content_index=0, part=part), + event("response.output_item.done", 6, output_index=0, item=item), + event("response.completed", 7, response=response("completed", [item])), + ] + + +def response_events(case): + if case == "success": + return message_events() + if case == "tool": + return tool_events() + if case == "upstream_failed": + failed = response("failed") + failed["error"] = {"code": "server_error", "message": ERROR_MESSAGE} + return [event("response.failed", 11, response=failed)] + code = "server_error" if case in ("error500", "after_tool") else "rate_limit_exceeded" + if case == "numeric429": + code = "429" + error_type = "server_error" if code == "server_error" else "rate_limit_error" + if case == "numeric429": + error_type = None + failure = {"type": "error", "error": {"message": ERROR_MESSAGE, + "type": error_type, "code": code, "param": "input"}} + prefix = [] + if case == "after_created": + prefix = [event("response.created", 4, response=response())] + elif case == "after_tool": + prefix = tool_events()[:3] + return [*prefix, failure] + + +def chat_events(): + chunk = {"id": "chatcmpl-fixture", "object": "chat.completion.chunk", + "created": 1, "model": MODEL} + return [ + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", + "content": SUCCESS}, "finish_reason": None}]}, + {**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + ] + + +class FixtureHandler(http.server.BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self): + payload = json.loads(self.rfile.read(int(self.headers["Content-Length"]))) + match = re.search(r"PROBE_CASE:([a-z0-9_]+)", json.dumps(payload)) + case = match.group(1) if match else "success" + self.server.requests_seen.append((self.path, case)) + if self.path.endswith("/chat/completions"): + chunks = chat_events() + if case == "chat_error": + chunks = chunks[:1] + [{"error": {"message": ERROR_MESSAGE, + "type": "server_error", "code": "server_error"}}] + body = "".join("data: " + json.dumps(chunk) + "\n\n" for chunk in chunks) + body += "data: [DONE]\n\n" + elif self.path.endswith("/responses"): + body = "".join("event: " + chunk["type"] + "\ndata: " + + json.dumps(chunk) + "\n\n" for chunk in response_events(case)) + else: + self.send_error(404, "Unexpected synthetic upstream path") + return + encoded = body.encode() + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Content-Length", str(len(encoded))) + self.send_header("Connection", "close") + self.end_headers() + self.wfile.write(encoded) + self.wfile.flush() + self.close_connection = True + + def log_message(self, *args): + pass + + +def parse_sse(body): + events = [] + for frame in body.replace("\r\n", "\n").split("\n\n"): + name = None + data = [] + for line in frame.splitlines(): + if line.startswith("event:"): + name = line[6:].strip() + elif line.startswith("data:"): + data.append(line[5:].lstrip()) + if data and data != ["[DONE]"]: + events.append((name, json.loads("\n".join(data)))) + return events + + +def post_raw(gateway_url, path, case, chat=False): + payload = {"model": MODEL, "stream": True} + prompt = "PROBE_CASE:" + case + if chat: + payload["messages"] = [{"role": "user", "content": prompt}] + else: + payload["input"] = prompt + request = urllib.request.Request(gateway_url + path, json.dumps(payload).encode(), + {"Content-Type": "application/json"}) + try: + with urllib.request.urlopen(request, timeout=30) as result: + return result.status, result.headers.get_content_type(), result.read().decode() + except urllib.error.HTTPError as error: + return error.code, error.headers.get_content_type(), error.read().decode() + + +def post(gateway_url, path, case, chat=False): + status, content_type, body = post_raw(gateway_url, path, case, chat=chat) + assert status == 200, (status, body) + assert content_type == "text/event-stream", (content_type, body) + return parse_sse(body) + + +def check_failure(gateway_url, case, expected_code, prefix_types): + frames = post(gateway_url, "/v1/responses", case) + assert frames, f"{case}: empty SSE stream" + name, failure = frames[-1] + assert name == "response.failed", (case, frames) + assert failure["type"] == "response.failed", failure + failed_response = failure["response"] + assert failed_response["object"] == "response", failed_response + assert failed_response["status"] == "failed", failed_response + assert failed_response["id"], failed_response + assert ERROR_MESSAGE in failed_response["error"]["message"], failure + assert failed_response["error"]["code"] == expected_code, failure + prefix = [data for _, data in frames[:-1]] + assert [data["type"] for data in prefix] == prefix_types, frames + assert failure["sequence_number"] > max( + (data.get("sequence_number", -1) for data in prefix), default=-1 + ), frames + if prefix: + assert failed_response["id"] == prefix[0]["response"]["id"], frames + if case == "after_tool": + assert prefix[-1]["delta"] == TOOL_ARGUMENTS, frames + if case == "upstream_failed": + assert failure["sequence_number"] == 11, frames + + +def check_success(gateway_url, case): + events = [data for _, data in post(gateway_url, "/v1/responses", case)] + expected = response_events(case) + assert [data["type"] for data in events] == [data["type"] for data in expected], events + assert [data["sequence_number"] for data in events] == list(range(len(events))), events + completed = events[-1]["response"] + assert completed["status"] == "completed", completed + if case == "tool": + assert events[2]["delta"] == TOOL_ARGUMENTS, events + assert completed["output"][0]["arguments"] == TOOL_ARGUMENTS, completed + else: + assert events[3]["delta"] == SUCCESS, events + assert completed["output"][0]["content"][0]["text"] == SUCCESS, completed + + +def check_chat(gateway_url, path="/v1/chat/completions", response_input=False): + frames = post(gateway_url, path, "success", chat=not response_input) + chunks = [data for _, data in frames] + assert chunks and all(chunk["object"] == "chat.completion.chunk" for chunk in chunks), chunks + assert all(name != "response.failed" for name, _ in frames), frames + text = "".join(choice["delta"].get("content", "") + for chunk in chunks for choice in chunk["choices"]) + assert text == SUCCESS, chunks + assert chunks[-1]["choices"][0]["finish_reason"] == "stop", chunks + + +def check_legacy_failure(gateway_url, path, case, chat): + frames = post(gateway_url, path, case, chat=chat) + assert frames, "Empty legacy error stream" + assert all(name != "response.failed" and data.get("type") != "response.failed" + for name, data in frames), frames + name, failure = frames[-1] + assert name is None, frames + assert ERROR_MESSAGE in failure["error"]["message"], frames + + +def check_cursor_pre_stream_failure(gateway_url): + status, content_type, body = post_raw(gateway_url, "/cursor/chat/completions", "error500") + assert status == 500, (status, body) + assert content_type == "application/json", (content_type, body) + failure = json.loads(body) + assert set(failure) == {"error"}, failure + assert failure["error"]["code"] == "500", failure + assert ERROR_MESSAGE in failure["error"]["message"], failure + + +def run_http_checks(gateway_url): + cases = [ + ("error429", "rate_limit_exceeded", []), + ("error500", "server_error", []), + ("numeric429", "rate_limit_exceeded", []), + ("upstream_failed", "server_error", []), + ("after_created", "rate_limit_exceeded", ["response.created"]), + ("after_tool", "server_error", ["response.created", "response.output_item.added", + "response.function_call_arguments.delta"]), + ] + for case, code, prefix in cases: + check_failure(gateway_url, case, code, prefix) + print(f"PASS Responses {case}", flush=True) + for case in ("success", "tool"): + check_success(gateway_url, case) + print(f"PASS Responses {case}", flush=True) + check_chat(gateway_url) + print("PASS Chat Completions success", flush=True) + check_legacy_failure(gateway_url, "/v1/chat/completions", "chat_error", chat=True) + print("PASS Chat Completions streaming error", flush=True) + check_chat(gateway_url, "/cursor/chat/completions", response_input=True) + print("PASS Cursor Responses bridge success", flush=True) + check_cursor_pre_stream_failure(gateway_url) + print("PASS Cursor Responses bridge pre-stream JSON error", flush=True) + check_legacy_failure(gateway_url, "/cursor/chat/completions", "after_tool", chat=False) + print("PASS Cursor Responses bridge streaming error after tool delta", flush=True) + + +def free_port(): + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + return sock.getsockname()[1] + + +def wait_ready(process, url, log_path): + deadline = time.monotonic() + 90 + while process.poll() is None and time.monotonic() < deadline: + try: + with urllib.request.urlopen(url + "/health/readiness", timeout=2) as result: + if result.status == 200: + return + except (OSError, urllib.error.URLError): + time.sleep(0.2) + raise RuntimeError("Synthetic proxy failed to start:\n" + log_path.read_text()[-12000:]) + + +@contextlib.contextmanager +def synthetic_proxy(port, host): + binary = shutil.which("litellm") + if not binary: + raise RuntimeError("Run this check in the patched LiteLLM image; litellm CLI is missing") + upstream = http.server.ThreadingHTTPServer(("127.0.0.1", 0), FixtureHandler) + upstream.requests_seen = [] + threading.Thread(target=upstream.serve_forever, daemon=True).start() + process = None + with tempfile.TemporaryDirectory(prefix="responses-stream-test-") as directory: + root = Path(directory) + config = {"model_list": [{"model_name": MODEL, "litellm_params": { + "model": "openai/" + MODEL, + "api_base": f"http://127.0.0.1:{upstream.server_port}/v1", + "api_key": "unused-fixture-key"}}], + "router_settings": {"num_retries": 0}, + "litellm_settings": {"num_retries": 0, "request_timeout": 15}} + config_path = root / "config.json" + config_path.write_text(json.dumps(config)) + log_path = root / "proxy.log" + selected_port = port or free_port() + url = f"http://127.0.0.1:{selected_port}" + env = {key: value for key, value in os.environ.items() + if key not in ("DATABASE_URL", "DIRECT_URL", "LITELLM_MASTER_KEY")} + try: + with log_path.open("w") as log: + process = subprocess.Popen([binary, "--config", str(config_path), + "--port", str(selected_port), "--host", host], + env=env, stdout=log, stderr=subprocess.STDOUT) + wait_ready(process, url, log_path) + yield url, upstream + except BaseException: + print(log_path.read_text()[-12000:], flush=True) + raise + finally: + if process is not None and process.poll() is None: + process.terminate() + try: + process.wait(timeout=10) + except subprocess.TimeoutExpired: + process.kill() + process.wait() + upstream.shutdown() + upstream.server_close() + + +def run_codex_checks(url, binary): + parsed = urllib.parse.urlparse(url) + if parsed.scheme != "http" or parsed.hostname not in ("127.0.0.1", "localhost", "::1"): + raise ValueError("--codex-url must point to the loopback synthetic proxy") + with tempfile.TemporaryDirectory(prefix="codex-stream-test-") as directory: + for case in ("error429", "error500", "success"): + output = Path(directory) / (case + ".txt") + provider = ('{name="Synthetic integration",base_url=' + json.dumps(url + "/v1") + + ',wire_api="responses",requires_openai_auth=false,' + 'request_max_retries=0,stream_max_retries=0,stream_idle_timeout_ms=10000}') + command = [binary, "exec", "--ignore-user-config", "--ignore-rules", "--ephemeral", + "--skip-git-repo-check", "--json", "--sandbox", "read-only", + "-C", directory, "--output-last-message", str(output), + "-c", 'model_provider="synthetic_test"', + "-c", "model_providers.synthetic_test=" + provider, + "-c", 'model_reasoning_effort="low"', + "-m", MODEL, "PROBE_CASE:" + case] + result = subprocess.run(command, capture_output=True, text=True, timeout=40) + combined = result.stdout + result.stderr + if case == "success": + assert result.returncode == 0, combined + assert output.exists() and SUCCESS in output.read_text(), combined + else: + assert result.returncode != 0, combined + assert ERROR_MESSAGE in combined, combined + assert "stream closed before response.completed" not in combined, combined + print(f"PASS Codex {case}", flush=True) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--serve", action="store_true", help="Keep the tested synthetic proxy running") + parser.add_argument("--port", type=int, default=0) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--codex-url", help="Test a local port-forward to the synthetic proxy") + parser.add_argument("--codex-bin", default="codex") + args = parser.parse_args() + if args.codex_url: + run_codex_checks(args.codex_url.rstrip("/"), args.codex_bin) + return + with synthetic_proxy(args.port, args.host) as (url, upstream): + run_http_checks(url) + assert upstream.requests_seen, "No requests reached the synthetic upstream" + print(json.dumps({"status": "ready", "gateway_url": url, + "synthetic_requests": len(upstream.requests_seen)}), flush=True) + if args.serve: + threading.Event().wait() + + +if __name__ == "__main__": + main() diff --git a/deploy/litellm/tests/test_responses_stream_errors.py b/deploy/litellm/tests/test_responses_stream_errors.py new file mode 100644 index 00000000..812387c3 --- /dev/null +++ b/deploy/litellm/tests/test_responses_stream_errors.py @@ -0,0 +1,178 @@ +# Copyright Advanced Micro Devices, Inc. +# SPDX-License-Identifier: MIT + +import copy +import importlib.util +import json +from pathlib import Path +import tempfile +from types import SimpleNamespace +import unittest + + +PATCHES = Path(__file__).resolve().parents[1] / "patches" + + +def load_module(name): + spec = importlib.util.spec_from_file_location(name, PATCHES / (name + ".py")) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +helper = load_module("responses_stream_errors") +installer = load_module("apply_responses_stream_errors") + + +class UpstreamError(Exception): + def __init__(self, status_code, message="Synthetic upstream failure"): + super().__init__(message) + self.status_code = status_code + self.message = message + + +def decode_failure(frame): + assert frame.startswith("event: response.failed\ndata: "), frame + assert frame.endswith("\n\n"), frame + return json.loads(frame.splitlines()[1][6:]) + + +class ResponsesStreamErrorTests(unittest.TestCase): + def test_rate_limit_and_service_failure_are_responses_events(self): + for status, code in ((429, "rate_limit_exceeded"), (500, "server_error")): + with self.subTest(status=status): + state = helper.ResponsesStreamErrorState() + failure = decode_failure(state.format_failure(UpstreamError(status))) + self.assertEqual(failure["type"], "response.failed") + self.assertEqual(failure["sequence_number"], 0) + response = failure["response"] + self.assertTrue(response["id"].startswith("resp_")) + self.assertEqual(response["object"], "response") + self.assertEqual(response["status"], "failed") + self.assertEqual(response["output"], []) + self.assertEqual(response["error"], { + "code": code, "message": "Synthetic upstream failure"}) + + def test_created_response_identity_and_next_sequence_are_preserved(self): + state = helper.ResponsesStreamErrorState() + state.observe_chunk(SimpleNamespace(type="response.created", sequence_number=8, + response=SimpleNamespace(id="resp_existing", model="gpt-5.2", + created_at=123))) + state.observe_chunk({"type": "response.output_text.delta", "sequence_number": 3}) + failure = decode_failure(state.format_failure(UpstreamError(500))) + self.assertEqual(failure["sequence_number"], 9) + self.assertEqual(failure["response"]["id"], "resp_existing") + self.assertEqual(failure["response"]["model"], "gpt-5.2") + self.assertEqual(failure["response"]["created_at"], 123) + + def test_tool_chunk_is_not_mutated_or_replayed_as_output(self): + state = helper.ResponsesStreamErrorState() + chunk = {"type": "response.function_call_arguments.delta", "sequence_number": 6, + "response_id": "resp_tools", "delta": '{"value":"partial'} + before = copy.deepcopy(chunk) + state.observe_chunk(chunk) + failure = decode_failure(state.format_failure(UpstreamError(500))) + self.assertEqual(chunk, before) + self.assertEqual(failure["response"]["id"], "resp_tools") + self.assertEqual(failure["sequence_number"], 7) + self.assertEqual(failure["response"]["output"], []) + + def test_terminal_response_cannot_be_followed_by_a_second_terminal(self): + for terminal in ("response.completed", "response.failed", "response.incomplete"): + with self.subTest(terminal=terminal): + state = helper.ResponsesStreamErrorState() + state.observe_chunk({"type": terminal, "sequence_number": 7}) + self.assertIsNone(state.format_failure(UpstreamError(500))) + + def test_emitted_failure_is_terminal(self): + state = helper.ResponsesStreamErrorState() + self.assertIsNotNone(state.format_failure(UpstreamError(429))) + self.assertIsNone(state.format_failure(UpstreamError(500))) + + def test_failure_serializing_a_terminal_chunk_still_reports_an_error(self): + state = helper.ResponsesStreamErrorState() + state.observe_chunk({"type": "response.completed", "sequence_number": 7}, emitted=False) + failure = decode_failure(state.format_failure(UpstreamError(500, "Serialization failed"))) + self.assertEqual(failure["type"], "response.failed") + self.assertEqual(failure["response"]["error"]["message"], "Serialization failed") + + def test_successfully_serialized_terminal_chunk_prevents_a_later_failure(self): + state = helper.ResponsesStreamErrorState() + state.observe_chunk({"type": "response.completed", "sequence_number": 7}, emitted=False) + state.mark_emitted() + self.assertIsNone(state.format_failure(UpstreamError(500))) + + def test_error_metadata_survives_a_fallback_wrapper_without_changing_status(self): + upstream = UpstreamError(500, "Mapped transport exception") + error = {"message": "Original provider message", "code": "429", "param": "input"} + helper.preserve_upstream_error(upstream, error) + wrapper = Exception("Router fallback exhausted") + wrapper.original_exception = upstream + failure = decode_failure(helper.ResponsesStreamErrorState().format_failure(wrapper)) + self.assertEqual(upstream.status_code, 500) + self.assertEqual(failure["response"]["error"], { + "code": "rate_limit_exceeded", "message": "Original provider message"}) + + def test_numeric_and_named_rate_limits_are_not_reported_as_server_errors(self): + for code in (429, "429", "rate_limit_exceeded", "insufficient_quota"): + with self.subTest(code=code): + upstream = UpstreamError(500) + helper.preserve_upstream_error(upstream, {"code": code}) + failure = decode_failure(helper.ResponsesStreamErrorState().format_failure(upstream)) + expected = "insufficient_quota" if code == "insufficient_quota" else "rate_limit_exceeded" + self.assertEqual(failure["response"]["error"]["code"], expected) + self.assertEqual(upstream.status_code, 500) + + def test_failed_upstream_event_keeps_its_identity_when_no_chunk_was_emitted(self): + upstream = UpstreamError(500) + event = {"type": "response.failed", "sequence_number": 11, + "response": {"id": "resp_upstream", "created_at": 12}} + helper.preserve_upstream_error(upstream, {"code": "server_error"}, event) + failure = decode_failure(helper.ResponsesStreamErrorState().format_failure(upstream)) + self.assertEqual(failure["sequence_number"], 11) + self.assertEqual(failure["response"]["id"], "resp_upstream") + self.assertEqual(failure["response"]["created_at"], 12) + + def test_raw_upstream_id_cannot_replace_an_already_visible_response_id(self): + state = helper.ResponsesStreamErrorState() + state.observe_chunk({"type": "response.created", "sequence_number": 0, + "response": {"id": "resp_client_visible"}}) + upstream = UpstreamError(500) + event = {"type": "response.failed", "sequence_number": 11, + "response": {"id": "resp_upstream_raw"}} + helper.preserve_upstream_error(upstream, {"code": "server_error"}, event) + failure = decode_failure(state.format_failure(upstream)) + self.assertEqual(failure["response"]["id"], "resp_client_visible") + self.assertEqual(failure["sequence_number"], 11) + + def test_error_messages_cannot_break_sse_framing(self): + message = 'Line one\n\ndata: forged event\n"quoted" ☃' + failure = decode_failure(helper.ResponsesStreamErrorState().format_failure( + UpstreamError(500, message))) + self.assertEqual(failure["response"]["error"]["message"], message) + + def test_missing_response_ids_are_unique_between_requests(self): + first = decode_failure(helper.ResponsesStreamErrorState().format_failure(UpstreamError(500))) + second = decode_failure(helper.ResponsesStreamErrorState().format_failure(UpstreamError(500))) + self.assertNotEqual(first["response"]["id"], second["response"]["id"]) + + +class InstallerTests(unittest.TestCase): + def test_unknown_upstream_version_fails_before_writing_any_files(self): + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + originals = {} + for relative in installer.BASE_SHA256: + path = root / relative + path.parent.mkdir(parents=True, exist_ok=True) + originals[path] = b"# An unsupported upstream source version.\n" + path.write_bytes(originals[path]) + with self.assertRaisesRegex(ValueError, "Unsupported LiteLLM source"): + installer.apply_patch(root) + for path, content in originals.items(): + self.assertEqual(path.read_bytes(), content) + self.assertFalse((root / "proxy" / installer.HELPER_NAME).exists()) + + +if __name__ == "__main__": + unittest.main() diff --git a/scripts/release-tests/litellm-build-context.sh b/scripts/release-tests/litellm-build-context.sh new file mode 100755 index 00000000..0cdbaff3 --- /dev/null +++ b/scripts/release-tests/litellm-build-context.sh @@ -0,0 +1,64 @@ +#!/usr/bin/env bash +# Copyright Advanced Micro Devices, Inc. +# SPDX-License-Identifier: MIT + +set -euo pipefail + +repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +tmp="$(mktemp -d)" +trap 'rm -rf "$tmp"' EXIT +mkdir -p "$tmp/bin" + +cat >"$tmp/bin/kubectl" <<'EOF' +#!/usr/bin/env bash +set -euo pipefail +for arg in "$@"; do + case "$arg" in + --from-file=build-context.tar.gz=*) + cp "${arg#--from-file=build-context.tar.gz=}" "$BUILD_CONTEXT_CAPTURE" + printf 'apiVersion: v1\nkind: ConfigMap\n' + exit 0 + ;; + esac +done +case "$*" in + *"apply -f -"*) cat >>"$BUILD_JOB_CAPTURE" ;; +esac +EOF +chmod +x "$tmp/bin/kubectl" + +env -u HARBOR_PASSWORD "PATH=$tmp/bin:$PATH" \ + "BUILD_CONTEXT_CAPTURE=$tmp/context.tar.gz" "BUILD_JOB_CAPTURE=$tmp/job.yaml" \ + REGISTRY=registry.example.invalid/gateway PUSH_SECRET=build-fixture-credentials \ + NAMESPACE=build-fixture TAG=v1.99.0-build-test \ + bash "$repo_root/deploy/litellm/build.sh" >"$tmp/build.log" 2>&1 || { + cat "$tmp/build.log" >&2 + exit 1 + } + +python3 - "$tmp/context.tar.gz" <<'PY' +import sys +import tarfile + +required = { + "Dockerfile", + "apim_key_hook.py", + "patches/apply_responses_stream_errors.py", + "patches/responses_stream_errors.py", +} +with tarfile.open(sys.argv[1]) as archive: + names = set(archive.getnames()) + missing = required - names + if missing: + raise SystemExit(f"Missing build context files: {sorted(missing)}") + for name in names: + if "__pycache__" in name or name.endswith(".pyc"): + raise SystemExit(f"Unexpected bytecode in build context: {name}") + for name in required: + if name.endswith(".py"): + compile(archive.extractfile(name).read(), name, "exec") +PY + +grep -qF 'tar -xzf /cm/build-context.tar.gz -C /workspace' "$tmp/job.yaml" +grep -qF 'secretName: build-fixture-credentials' "$tmp/job.yaml" +echo "LiteLLM patched image build context: ok"