diff --git a/.circleci/config.yml b/.circleci/config.yml index 599c58a40e..e1872ff00e 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1625,6 +1625,25 @@ jobs: command: | zstd -d litellm-docker-database.tar.zst --stdout | docker load docker tag litellm-docker-database:ci my-app:latest + - run: + name: Start OpenAI image record/replay proxy + background: true + command: | + CASSETTE_REDIS_URL="$CASSETTE_REDIS_URL" \ + RECORDER_UPSTREAM_BASE_URL="https://api.openai.com" \ + uv run --no-sync python tests/_openai_record_replay_proxy.py --host 0.0.0.0 --port 8090 + - run: + name: Wait for record/replay proxy + command: | + for i in $(seq 1 30); do + if curl -sf http://localhost:8090/__recorder_health >/dev/null 2>&1; then + echo "record/replay proxy is up" + exit 0 + fi + sleep 1 + done + echo "record/replay proxy did not become ready" >&2 + exit 1 - run: name: Run Docker container command: | @@ -1655,6 +1674,7 @@ jobs: -e LANGFUSE_PROJECT2_PUBLIC=$LANGFUSE_PROJECT2_PUBLIC \ -e LANGFUSE_PROJECT1_SECRET=$LANGFUSE_PROJECT1_SECRET \ -e LANGFUSE_PROJECT2_SECRET=$LANGFUSE_PROJECT2_SECRET \ + -e IMAGE_GEN_RECORDER_BASE_URL=http://host.docker.internal:8090/v1 \ --add-host host.docker.internal:host-gateway \ --name my-app \ -v $(pwd)/proxy_server_config.yaml:/app/config.yaml \ diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index da37eb3428..7883a15e5a 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -44,6 +44,15 @@ model_list: - model_name: openai-dall-e-3 # dall-e-3 deprecated 2026-05-12; underlying now gpt-image-1 litellm_params: model: gpt-image-1 + # In CI, IMAGE_GEN_RECORDER_BASE_URL points this at the record/replay proxy + # (tests/_openai_record_replay_proxy.py) so the image spend E2E doesn't depend + # on OpenAI's uptime every commit. Unset elsewhere, so it resolves to None and + # falls back to api.openai.com. + - model_name: gpt-image-1 + litellm_params: + model: openai/gpt-image-1 + api_key: os.environ/OPENAI_API_KEY + api_base: os.environ/IMAGE_GEN_RECORDER_BASE_URL - model_name: fake-openai-endpoint litellm_params: model: openai/gpt-5-mini diff --git a/tests/_openai_record_replay_proxy.py b/tests/_openai_record_replay_proxy.py new file mode 100644 index 0000000000..60832188a7 --- /dev/null +++ b/tests/_openai_record_replay_proxy.py @@ -0,0 +1,238 @@ +"""Record/replay reverse proxy for the dockerized image-gen spend E2E. + +The spend-accuracy test ``tests/test_keys.py:: +test_key_info_spend_values_image_generation`` runs the litellm proxy in its own +container and curls it over real HTTP, then asserts the proxy tracked a nonzero +spend for a ``gpt-image-1`` call. That call wildcard-routes to ``openai/*`` on +the real key, so every commit run hit api.openai.com for a paid image and was +exposed to OpenAI outages (the 401 that started this). + +This process sits between the proxy and api.openai.com. The proxy points only +its image model's ``api_base`` here; nothing else about the topology changes. +The first request (or the first after a recording lapses) is forwarded live to +OpenAI and recorded; subsequent requests within the TTL replay the recorded +response, so the per-commit run no longer depends on OpenAI being up. + +Recordings live in the same Redis cassette store as the VCR persister +(``CASSETTE_REDIS_URL``) and expire ``CASSETTE_TTL_SECONDS`` after their last +write, never refreshed on read. A recording therefore goes stale a day after +capture and the next run past that point re-records live and catches provider +contract drift, exactly matching the lapse-after-write contract in +``tests/_vcr_redis_persister.py``. +""" + +from __future__ import annotations + +import base64 +import hashlib +import json +import os +from typing import Awaitable, Callable, List, Optional, Tuple + +CASSETTE_TTL_SECONDS = 24 * 60 * 60 +RECORD_KEY_PREFIX = "litellm:openai:record:" +RECORDER_REDIS_URL_ENV = "CASSETTE_REDIS_URL" +UPSTREAM_BASE_URL_ENV = "RECORDER_UPSTREAM_BASE_URL" +DEFAULT_UPSTREAM_BASE_URL = "https://api.openai.com" + +Headers = List[Tuple[str, str]] +UpstreamResult = Tuple[int, Headers, bytes] +FetchUpstream = Callable[[], Awaitable[UpstreamResult]] + +# Headers the re-serving layer owns and must set itself. Replaying an upstream +# framing header verbatim onto a freshly built response is the same class of bug +# as the Bedrock content-length: 0 regression (#29549): a stale header rides +# along and contradicts the real body. The serving server recomputes +# content-length and sets its own date/server; the stored body is already +# content-decoded so content-encoding must not claim otherwise. +_STRIPPED_RESPONSE_HEADERS = frozenset( + { + "content-length", + "content-encoding", + "transfer-encoding", + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailer", + "upgrade", + "date", + "server", + } +) + + +def _canonical_body(body: bytes) -> bytes: + if not body: + return b"" + try: + return json.dumps(json.loads(body), sort_keys=True, separators=(",", ":")).encode("utf-8") + except (ValueError, TypeError): + return body + + +def _sanitize_headers(headers: Headers) -> Headers: + return [(k, v) for (k, v) in headers if k.lower() not in _STRIPPED_RESPONSE_HEADERS] + + +class OpenAIRecordReplay: + """Record-once / replay-from-Redis for upstream OpenAI HTTP calls. + + ``redis_client`` is injected so the process wiring and the tests share one + code path; pass ``None`` to run as a pure live passthrough (local dev with + no cassette Redis). + """ + + def __init__( + self, + redis_client, + *, + upstream_base_url: str = DEFAULT_UPSTREAM_BASE_URL, + ttl_seconds: int = CASSETTE_TTL_SECONDS, + ) -> None: + self._redis = redis_client + self.upstream_base_url = upstream_base_url.rstrip("/") + self._ttl_seconds = ttl_seconds + + @staticmethod + def record_key(method: str, path: str, body: bytes) -> str: + digest = hashlib.sha256( + b"\n".join( + [ + method.upper().encode("utf-8"), + path.encode("utf-8"), + _canonical_body(body), + ] + ) + ).hexdigest() + return f"{RECORD_KEY_PREFIX}{digest}" + + async def handle(self, method: str, path: str, body: bytes, fetch_upstream: FetchUpstream) -> UpstreamResult: + key = self.record_key(method, path, body) + cached = self._cache_get(key) + if cached is not None: + return cached + + status, headers, resp_body = await fetch_upstream() + sanitized = _sanitize_headers(headers) + if 200 <= status < 300: + self._cache_set(key, status, sanitized, resp_body) + return status, sanitized, resp_body + + def _cache_get(self, key: str) -> Optional[UpstreamResult]: + if self._redis is None: + return None + try: + raw = self._redis.get(key) + except Exception: + return None + if raw is None: + return None + try: + payload = json.loads(raw) + status = int(payload["status"]) + headers = [(str(k), str(v)) for k, v in payload["headers"]] + resp_body = base64.b64decode(payload["body_b64"]) + except Exception: + return None + return status, headers, resp_body + + def _cache_set(self, key: str, status: int, headers: Headers, body: bytes) -> None: + if self._redis is None: + return + payload = json.dumps( + { + "status": status, + "headers": [[k, v] for (k, v) in headers], + "body_b64": base64.b64encode(body).decode("ascii"), + } + ) + try: + self._redis.set(key, payload, ex=self._ttl_seconds) + except Exception: + pass + + +def _build_default_redis_client(): + url = os.environ.get(RECORDER_REDIS_URL_ENV) + if not url: + return None + import redis + + return redis.Redis.from_url( + url, + socket_timeout=5, + socket_connect_timeout=5, + decode_responses=False, + ) + + +def create_app(recorder: Optional[OpenAIRecordReplay] = None, http_client=None): + import contextlib + + import httpx + from starlette.applications import Starlette + from starlette.responses import PlainTextResponse, Response + from starlette.routing import Route + + if recorder is None: + recorder = OpenAIRecordReplay( + redis_client=_build_default_redis_client(), + upstream_base_url=os.environ.get(UPSTREAM_BASE_URL_ENV, DEFAULT_UPSTREAM_BASE_URL), + ) + owns_client = http_client is None + client = http_client or httpx.AsyncClient(timeout=httpx.Timeout(120.0)) + + @contextlib.asynccontextmanager + async def lifespan(_app): + try: + yield + finally: + if owns_client: + await client.aclose() + + async def health(_request): + return PlainTextResponse("ok") + + async def proxy(request): + body = await request.body() + path = request.url.path + full_path = f"{path}?{request.url.query}" if request.url.query else path + + async def fetch_upstream() -> UpstreamResult: + fwd_headers = {k: v for k, v in request.headers.items() if k.lower() != "host"} + upstream = await client.request( + request.method, + f"{recorder.upstream_base_url}{full_path}", + content=body, + headers=fwd_headers, + ) + return ( + upstream.status_code, + list(upstream.headers.items()), + upstream.content, + ) + + status, headers, resp_body = await recorder.handle(request.method, full_path, body, fetch_upstream) + return Response(content=resp_body, status_code=status, headers=dict(headers)) + + return Starlette( + routes=[ + Route("/__recorder_health", health, methods=["GET"]), + Route("/{path:path}", proxy, methods=["GET", "POST", "PUT", "PATCH", "DELETE"]), + ], + lifespan=lifespan, + ) + + +if __name__ == "__main__": + import argparse + + import uvicorn + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--host", default="0.0.0.0") + parser.add_argument("--port", type=int, default=8090) + args = parser.parse_args() + uvicorn.run(create_app(), host=args.host, port=args.port) diff --git a/tests/llm_translation/test_openai_record_replay_proxy.py b/tests/llm_translation/test_openai_record_replay_proxy.py new file mode 100644 index 0000000000..5e34621d01 --- /dev/null +++ b/tests/llm_translation/test_openai_record_replay_proxy.py @@ -0,0 +1,216 @@ +from __future__ import annotations + +import asyncio +import os +import sys + +import fakeredis + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) + +from tests._openai_record_replay_proxy import ( # noqa: E402 + CASSETTE_TTL_SECONDS, + RECORD_KEY_PREFIX, + OpenAIRecordReplay, +) + +_OK_BODY = b'{"data":[{"b64_json":"aW1n"}],"usage":{"total_tokens":42}}' + + +class _Upstream: + """Stub live upstream; counts calls so replays can be proven offline.""" + + def __init__(self, status=200, headers=None, body=_OK_BODY): + self.calls = 0 + self._status = status + self._headers = headers if headers is not None else [("content-type", "application/json")] + self._body = body + + async def __call__(self): + self.calls += 1 + return self._status, list(self._headers), self._body + + +def _recorder(client=None): + return OpenAIRecordReplay(client if client is not None else fakeredis.FakeStrictRedis()) + + +def _run(coro): + return asyncio.run(coro) + + +def test_miss_forwards_to_upstream_and_records(): + fake = fakeredis.FakeStrictRedis() + recorder = _recorder(fake) + upstream = _Upstream() + + status, headers, body = _run( + recorder.handle("POST", "/v1/images/generations", b'{"model":"gpt-image-1"}', upstream) + ) + + assert upstream.calls == 1 + assert status == 200 + assert body == _OK_BODY + key = OpenAIRecordReplay.record_key("POST", "/v1/images/generations", b'{"model":"gpt-image-1"}') + assert key.startswith(RECORD_KEY_PREFIX) + assert fake.get(key) is not None + + +def test_hit_replays_without_calling_upstream(): + recorder = _recorder() + upstream = _Upstream() + body_in = b'{"model":"gpt-image-1","prompt":"otter"}' + + first = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + second = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + + assert upstream.calls == 1 + assert first == second + assert second[2] == _OK_BODY + + +def test_different_body_is_a_separate_recording(): + recorder = _recorder() + upstream = _Upstream() + + _run(recorder.handle("POST", "/v1/images/generations", b'{"prompt":"otter"}', upstream)) + _run(recorder.handle("POST", "/v1/images/generations", b'{"prompt":"seal"}', upstream)) + + assert upstream.calls == 2 + + +def test_record_key_ignores_json_key_order(): + a = OpenAIRecordReplay.record_key("POST", "/v1/images/generations", b'{"model":"x","prompt":"y"}') + b = OpenAIRecordReplay.record_key("POST", "/v1/images/generations", b'{"prompt":"y","model":"x"}') + assert a == b + + +def test_ttl_set_on_write_and_not_refreshed_on_read(): + """A replay must not slide the recording's expiry forward. + + The recording counts down from its last write so it lapses + ``CASSETTE_TTL_SECONDS`` after capture and the next run re-records live, + catching provider drift. Refreshing the TTL on a replay would keep an + actively-replayed recording alive forever and that drift check would never + run. This mirrors the VCR persister's lapse-after-write contract. + """ + fake = fakeredis.FakeStrictRedis() + recorder = _recorder(fake) + upstream = _Upstream() + body_in = b'{"model":"gpt-image-1"}' + key = OpenAIRecordReplay.record_key("POST", "/v1/images/generations", body_in) + + _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + assert CASSETTE_TTL_SECONDS - 5 <= fake.ttl(key) <= CASSETTE_TTL_SECONDS + + fake.expire(key, 60) + _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + + assert fake.ttl(key) <= 60 + + +def test_replay_drops_framing_headers_so_server_recomputes(): + fake = fakeredis.FakeStrictRedis() + recorder = _recorder(fake) + upstream = _Upstream( + headers=[ + ("content-type", "application/json"), + ("content-length", "9999"), + ("transfer-encoding", "chunked"), + ("content-encoding", "gzip"), + ("date", "Mon, 01 Jan 2024 00:00:00 GMT"), + ("server", "cloudflare"), + ("x-request-id", "req_abc"), + ] + ) + body_in = b'{"model":"gpt-image-1"}' + + _, live_headers, _ = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + _, replay_headers, _ = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + + for headers in (live_headers, replay_headers): + names = {k.lower() for k, _ in headers} + assert names.isdisjoint( + { + "content-length", + "transfer-encoding", + "content-encoding", + "date", + "server", + } + ) + assert ("content-type", "application/json") in headers + assert ("x-request-id", "req_abc") in headers + + +def test_non_2xx_response_is_not_cached(): + fake = fakeredis.FakeStrictRedis() + recorder = _recorder(fake) + upstream = _Upstream(status=500, body=b'{"error":"boom"}') + body_in = b'{"model":"gpt-image-1"}' + key = OpenAIRecordReplay.record_key("POST", "/v1/images/generations", body_in) + + status, _, _ = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + assert status == 500 + assert fake.get(key) is None + + _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + assert upstream.calls == 2 + + +class _BoomRedis: + def get(self, *args, **kwargs): + raise ConnectionError("redis offline") + + def set(self, *args, **kwargs): + raise ConnectionError("redis offline") + + +def test_redis_outage_degrades_to_live_passthrough(): + recorder = _recorder(_BoomRedis()) + upstream = _Upstream() + body_in = b'{"model":"gpt-image-1"}' + + first = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + second = _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + + assert first[0] == 200 and second[0] == 200 + assert upstream.calls == 2 + + +def test_passthrough_when_no_redis_client_configured(): + recorder = OpenAIRecordReplay(None) + upstream = _Upstream() + body_in = b'{"model":"gpt-image-1"}' + + _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + _run(recorder.handle("POST", "/v1/images/generations", body_in, upstream)) + + assert upstream.calls == 2 + + +class _StubClient: + def __init__(self): + self.closed = False + + async def aclose(self): + self.closed = True + + +def test_app_lifespan_leaves_injected_http_client_open(): + """The app must only close the client it created, never a caller's. + + A caller that injects its own client owns that client's lifecycle; the + app closing it would break reuse across multiple ``create_app`` calls. + """ + from starlette.testclient import TestClient + + from tests._openai_record_replay_proxy import create_app + + client = _StubClient() + app = create_app(recorder=_recorder(), http_client=client) + + with TestClient(app): + pass + + assert client.closed is False