headroom/tests/test_ws_http_fallback.py
Abhay Singh 536c949a69
fix(proxy/openai): propagate provider usage on the Responses WS->HTTP fallback (#2988)
## Description

When Codex uses the OpenAI Responses WebSocket endpoint through Headroom
and the upstream WebSocket is rejected, Headroom falls back to HTTPS
POST/SSE. On that fallback the dashboard reported zero or tiny input
tokens for a large request, and invalid savings:

```json
{ "input_tokens_original": 3, "input_tokens_optimized": 0,
  "output_tokens": 246, "tokens_saved": 31052, "savings_percent": 33233.33 }
```

## Root cause

`_ws_http_fallback` (openai.py) relays the SSE `data:` events to the
client but never parses the terminal `response.completed` event for
usage. The non-fallback WS path accumulates
`_extract_responses_usage(event)` into the session totals on every
`response.completed` frame (openai.py ~8182); the fallback path did not.
So `ws_input_tokens_total` stayed at the small local count, and the
session-end RequestLog computed `optimized_tokens =
residual_input_tokens = 0`, leaving `tokens_saved >
input_tokens_original` and `savings_percent` far above 100%.

## Fix

`_ws_http_fallback` now parses each relayed `response.completed` line
with the existing `_extract_responses_usage` and returns the accumulated
`(input, output, cache_read, cache_write, uncached)` provider usage. The
caller folds it into the WS session totals, so the session-end outcome
uses the authoritative provider wire-token count -- bringing the
fallback to parity with the non-fallback WS path. SSE relay behaviour is
otherwise unchanged.

Fixes #2957

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)
- [ ] New feature
- [ ] Breaking change
- [ ] Documentation update
- [ ] Performance improvement
- [ ] Code refactoring (no functional changes)

## Changes Made

- `headroom/proxy/handlers/openai.py` (`_ws_http_fallback`): accumulate
usage from `response.completed` SSE lines (both the main relay loop and
the buffer flush) and return the `(input, output, cache_read,
cache_write, uncached)` tuple from every exit path; the WS handler
caller adds it to `ws_input_tokens_total` / `ws_output_tokens_total` /
cache / uncached totals before the session-end RequestLog.
- `tests/test_ws_http_fallback.py`: the fallback returns the provider
usage from a `response.completed` event
(input/output/cache_read/uncached), and returns all-zeros when no
completed event arrives.

## Testing

- [x] Unit tests pass (`pytest`)
- [x] Linting passes (`ruff check`)
- [x] Type checking passes (`mypy`)
- [x] New tests added

### Test Output

```text
tests/test_ws_http_fallback.py  13 passed  (11 existing + 2 new)
# uvx ruff@0.15.22 check  -> All checks passed!
# uvx mypy@1.20.2 headroom/proxy/handlers/openai.py -> Success: no issues found in 1 source file
```

## Real Behavior Proof

- Environment: Windows 11, Python 3.12.11, project venv, pytest 9.1.1,
ruff 0.15.22 and mypy 1.20.2 via uvx.
- Exact command / steps: drove `_ws_http_fallback` with the existing
WS/stream mocks, feeding an SSE `response.completed` carrying
`usage.input_tokens=31055`, `output_tokens=246`,
`input_tokens_details.cached_tokens=20000`. The method now returns
`(31055, 246, 20000, ..., 11055)`; a stream with no completed event
returns all zeros. The existing 11 relay/routing/retry tests are
unchanged (they ignore the new return value).
- Observed result: the fallback surfaces the provider's real input
usage, so the WS session-end outcome records the actual input tokens
instead of 0, and savings percentages stay within a meaningful range.
- Not tested: a live Codex WS session that triggers the upstream-WS
rejection and HTTP fallback end to end (needs a real upstream refusing
the WS). The usage-propagation contract is verified at the fallback
boundary with the same mocks the existing fallback tests use.

## Runtime Rollout Safety

- Rollout-managed feature(s): none. The OpenAI Responses WS-to-HTTP
fallback is always-on transport behavior, not rollout-channel-gated.
- Minimum rollout channel: N/A (no rollout-managed behavior).
- Stable/default behavior changed: yes, as a bug fix. On the WS-to-HTTP
fallback the session-end outcome now records the provider's real
input/output/cache usage from `response.completed` instead of leaving
`ws_input_tokens_total` at 0 (which produced >100% savings). SSE relay
to the client is unchanged.
- Kill switch / disable path: N/A. This corrects accounting only; there
is no behavioral toggle and no user-facing surface beyond the recorded
outcome numbers.
- Unsafe override required: no.
- Qualification impact: fallback-path token accounting now matches the
non-fallback WS path and the HTTP Responses path (all three use
`_extract_responses_usage`); savings percentages return to a valid
range.
- Rollback path: revert this PR; the fallback returns to reporting zero
input usage on this path.

## Review Readiness

- [x] I have performed a self-review
- [x] This PR is ready for human review

## Checklist

- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [x] I have added tests that prove my fix is effective
- [x] New and existing unit tests pass locally with my changes
- [x] I did **not** edit `CHANGELOG.md`: it is generated by
release-please from my Conventional Commit PR title

## Additional Notes

The fix reuses the already-present `_extract_responses_usage` (same
parser the non-fallback WS path and HTTP Responses path use), so
cache-read/write and uncached accounting stay consistent across all
three transports.

Co-authored-by: JD Davis <mxjerrett@gmail.com>
2026-08-16 15:04:39 -07:00

392 lines
14 KiB
Python

"""Tests for WebSocket HTTP fallback in the OpenAI handler.
When the upstream WebSocket connection to OpenAI fails (HTTP 500),
the proxy should transparently fall back to HTTP POST streaming
and relay SSE events over the client WebSocket.
"""
from __future__ import annotations
import asyncio
import json
from types import SimpleNamespace
import httpx
class FakeWebSocket:
"""Minimal WebSocket mock for testing."""
def __init__(self):
self.sent_texts: list[str] = []
self.closed = False
async def send_text(self, data: str) -> None:
self.sent_texts.append(data)
async def close(self, code: int = 1000, reason: str = "") -> None:
self.closed = True
class FakeStreamResponse:
"""Mock httpx streaming response."""
def __init__(
self,
status_code: int = 200,
sse_events: list[str] | None = None,
headers: dict[str, str] | None = None,
):
self.status_code = status_code
self._events = sse_events or []
self.headers = headers or {}
async def aiter_text(self):
for event in self._events:
yield event
async def aiter_bytes(self):
yield b"error body"
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
class FakeHttpClient:
"""Mock httpx.AsyncClient with stream support."""
def __init__(self, response: FakeStreamResponse):
self._response = response
def stream(self, method, url, **kwargs):
return self._response
def _make_handler():
"""Create a minimal OpenAIHandlerMixin-like object."""
from headroom.proxy.handlers.openai import OpenAIHandlerMixin
obj = object.__new__(OpenAIHandlerMixin)
obj.OPENAI_API_URL = "https://api.openai.com"
obj.http_client = None
obj.config = SimpleNamespace(
retry_max_attempts=3,
retry_base_delay_ms=0,
retry_max_delay_ms=0,
)
return obj
class TestWsHttpFallback:
def test_fallback_relays_sse_events(self):
"""HTTP fallback should relay SSE data lines as WS text messages."""
handler = _make_handler()
ws = FakeWebSocket()
sse_lines = [
'event: response.created\ndata: {"type":"response.created","response":{"id":"r1"}}\n\n',
'event: response.output_item.added\ndata: {"type":"response.output_item.added"}\n\n',
'event: response.completed\ndata: {"type":"response.completed"}\n\n',
"data: [DONE]\n\n",
]
response = FakeStreamResponse(200, sse_lines)
handler.http_client = FakeHttpClient(response)
body = {"model": "gpt-5.4", "input": "hi"}
first_msg_raw = json.dumps({"type": "response.create", "response": body})
asyncio.run(
handler._ws_http_fallback(
ws, body, first_msg_raw, {"Authorization": "Bearer test"}, "req_1"
)
)
assert len(ws.sent_texts) == 3 # 3 data events, [DONE] skipped
assert '"response.created"' in ws.sent_texts[0]
assert '"response.output_item.added"' in ws.sent_texts[1]
assert '"response.completed"' in ws.sent_texts[2]
assert ws.closed
def test_fallback_sends_error_on_non_200(self):
"""HTTP fallback should send error event on non-200 response."""
handler = _make_handler()
ws = FakeWebSocket()
response = FakeStreamResponse(status_code=401)
handler.http_client = FakeHttpClient(response)
body = {"model": "gpt-5.4", "input": "hi"}
asyncio.run(
handler._ws_http_fallback(
ws, body, json.dumps(body), {"Authorization": "Bearer bad"}, "req_2"
)
)
assert len(ws.sent_texts) == 1
event = json.loads(ws.sent_texts[0])
assert event["type"] == "error"
assert "401" in event["error"]["message"]
def test_fallback_sets_stream_true(self):
"""HTTP fallback should force stream=True in request body.
After PR-A3 (byte-faithful Python forwarders) the fallback sends
the request body as raw bytes via `content=`, not via the `json=`
kwarg. The test extracts the posted JSON from the captured bytes.
"""
handler = _make_handler()
ws = FakeWebSocket()
captured_kwargs: dict = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_kwargs.update(kwargs)
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
body = {"model": "gpt-5.4", "input": "test", "stream": False}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_3"))
posted = json.loads(captured_kwargs["content"])
assert posted["stream"] is True
def test_fallback_unwraps_response_create_envelope(self):
"""HTTP fallback should unwrap WS response.create wrapper for HTTP POST."""
handler = _make_handler()
ws = FakeWebSocket()
captured_kwargs: dict = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_kwargs.update(kwargs)
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
# WS sends wrapped format: {"type": "response.create", "response": {...}}
inner = {
"model": "gpt-5.4",
"input": [{"role": "user", "content": [{"type": "input_text", "text": "hi"}]}],
}
ws_msg = {"type": "response.create", "response": inner}
asyncio.run(handler._ws_http_fallback(ws, ws_msg, json.dumps(ws_msg), {}, "req_unwrap"))
posted = json.loads(captured_kwargs["content"])
# Should be the inner response, not the wrapper
assert "type" not in posted # no "response.create" type field
assert posted["model"] == "gpt-5.4"
assert posted["stream"] is True
assert "input" in posted
def test_fallback_strips_top_level_response_create_type(self):
"""HTTP fallback should strip top-level response.create metadata."""
handler = _make_handler()
ws = FakeWebSocket()
captured_kwargs: dict = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_kwargs.update(kwargs)
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
body = {"type": "response.create", "model": "gpt-5.4", "input": "hi"}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_type_strip"))
posted = json.loads(captured_kwargs["content"])
assert posted["model"] == "gpt-5.4"
assert posted["stream"] is True
assert "type" not in posted
def test_fallback_handles_http_exception(self):
"""HTTP fallback should send error event when HTTP request fails."""
handler = _make_handler()
ws = FakeWebSocket()
class FailingClient:
def stream(self, method, url, **kwargs):
raise ConnectionError("upstream unreachable")
handler.http_client = FailingClient()
body = {"model": "gpt-5.4", "input": "test"}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_4"))
assert len(ws.sent_texts) == 1
event = json.loads(ws.sent_texts[0])
assert event["type"] == "error"
assert "unreachable" in event["error"]["message"]
def test_fallback_retries_connect_timeout(self):
"""HTTP fallback should retry transient connect timeouts."""
handler = _make_handler()
ws = FakeWebSocket()
attempts = {"count": 0}
class FlakyClient:
def stream(self, method, url, **kwargs):
attempts["count"] += 1
if attempts["count"] == 1:
raise httpx.ConnectTimeout("timed out")
return FakeStreamResponse(200, ['data: {"type":"response.completed"}\n\n'])
handler.http_client = FlakyClient()
body = {"model": "gpt-5.4", "input": "test"}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_retry"))
assert attempts["count"] == 2
assert len(ws.sent_texts) == 1
assert json.loads(ws.sent_texts[0])["type"] == "response.completed"
def test_fallback_routes_chatgpt_auth_to_chatgpt_domain(self):
"""ChatGPT session auth should route to chatgpt.com, not api.openai.com."""
handler = _make_handler()
ws = FakeWebSocket()
captured_url = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_url["url"] = url
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
body = {"model": "gpt-5.4", "input": "test"}
# ChatGPT session auth includes this header
headers = {
"Authorization": "Bearer chatgpt-session-token",
"ChatGPT-Account-ID": "acct_abc123",
}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), headers, "req_5"))
assert "chatgpt.com" in captured_url["url"]
assert "api.openai.com" not in captured_url["url"]
def test_fallback_chatgpt_auth_forces_store_false(self):
"""ChatGPT Responses backend requires explicit store=false."""
handler = _make_handler()
ws = FakeWebSocket()
captured_kwargs: dict = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_kwargs.update(kwargs)
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
body = {"model": "gpt-5.4", "input": "test", "store": True}
headers = {
"Authorization": "Bearer chatgpt-session-token",
"ChatGPT-Account-ID": "acct_abc123",
}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), headers, "req_store"))
posted = json.loads(captured_kwargs["content"])
assert posted["store"] is False
assert posted["stream"] is True
def test_fallback_routes_api_key_to_openai(self):
"""API key auth should route to api.openai.com."""
handler = _make_handler()
ws = FakeWebSocket()
captured_url = {}
class CapturingClient:
def stream(self, method, url, **kwargs):
captured_url["url"] = url
return FakeStreamResponse(200, ["data: [DONE]\n\n"])
handler.http_client = CapturingClient()
body = {"model": "gpt-5.4", "input": "test"}
# API key auth — no ChatGPT-Account-ID header
headers = {"Authorization": "Bearer sk-abc123"}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), headers, "req_6"))
assert "api.openai.com" in captured_url["url"]
def test_fallback_returns_provider_usage_from_completed_event(self):
"""The fallback must surface the provider's input usage (#2957).
Otherwise the WS session-end outcome records input_tokens=0 for a large
request and savings percentages blow past 100.
"""
handler = _make_handler()
ws = FakeWebSocket()
completed = {
"type": "response.completed",
"response": {
"usage": {
"input_tokens": 31055,
"output_tokens": 246,
"input_tokens_details": {"cached_tokens": 20000},
}
},
}
sse_lines = [
'data: {"type":"response.created","response":{"id":"r1"}}\n\n',
f"data: {json.dumps(completed)}\n\n",
"data: [DONE]\n\n",
]
handler.http_client = FakeHttpClient(FakeStreamResponse(200, sse_lines))
body = {"model": "gpt-5.4", "input": "big context"}
usage = asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_usage"))
input_tokens, output_tokens, cache_read, _cache_write, uncached = usage
assert input_tokens == 31055
assert output_tokens == 246
assert cache_read == 20000
assert uncached == 31055 - 20000
def test_fallback_returns_zero_usage_without_completed_event(self):
handler = _make_handler()
ws = FakeWebSocket()
handler.http_client = FakeHttpClient(
FakeStreamResponse(200, ['data: {"type":"response.created"}\n\n', "data: [DONE]\n\n"])
)
usage = asyncio.run(
handler._ws_http_fallback(
ws, {"model": "gpt-5.4", "input": "hi"}, json.dumps({"input": "hi"}), {}, "req_none"
)
)
assert usage == (0, 0, 0, 0, 0)
def test_fallback_refreshes_codex_rate_limit_state(self, monkeypatch):
"""A successful fallback refreshes Codex /stats from response headers.
The fallback can't forward headers onto the (already-accepted) client
101, but it should still keep Python /stats in sync so the gauge does
not go stale when the WS upgrade fails and we drop to HTTP.
"""
handler = _make_handler()
ws = FakeWebSocket()
captured: dict[str, dict[str, str]] = {}
class _FakeState:
def update_from_headers(self, hdrs):
captured["headers"] = dict(hdrs)
import headroom.subscription.codex_rate_limits as crl
monkeypatch.setattr(crl, "get_codex_rate_limit_state", lambda: _FakeState())
response = FakeStreamResponse(
200,
['data: {"type":"response.completed"}\n\n', "data: [DONE]\n\n"],
headers={
"x-codex-primary-used-percent": "42",
"content-type": "text/event-stream",
},
)
handler.http_client = FakeHttpClient(response)
body = {"model": "gpt-5.4", "input": "hi"}
asyncio.run(handler._ws_http_fallback(ws, body, json.dumps(body), {}, "req_capture"))
assert captured["headers"]["x-codex-primary-used-percent"] == "42"