headroom/tests/test_openai_codex_routing.py
John Xu 1c50eca8b3
fix(proxy): skip Responses memory tools for ChatGPT auth (#1579)
## Description

Fix ChatGPT/Codex session-auth Responses proxy handling so the ChatGPT
backend always receives an explicit `store=false`, while keeping
Responses memory tools limited to the regular API-key path where stored
responses are supported.

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)
- [ ] New feature (non-breaking change that adds functionality)
- [ ] Breaking change (fix or feature that would cause existing
functionality to change)
- [ ] Documentation update
- [ ] Performance improvement
- [ ] Code refactoring (no functional changes)

## Changes Made

- Detect ChatGPT auth before Responses memory-tool injection and force
`store=false` for ChatGPT-auth Responses payloads.
- Skip Responses memory tools and transparent memory-tool continuation
handling for ChatGPT auth across HTTP, WebSocket first frames, WebSocket
follow-up `response.create` frames, and WS-to-HTTP fallback.
- Preserve API-key behavior after the current main merge: API-key
requests that explicitly set `store=false` skip Responses memory tools,
while API-key requests that receive injected memory tools are forced to
`store=true` for continuation support.
- Address Copilot formatter comments by making
`_allow_responses_memory_tools` call sites formatter-stable.

## Testing

- [x] Unit tests pass (`pytest`)
- [x] Linting passes (`ruff check .`)
- [ ] Type checking passes (`mypy headroom`)
- [x] New tests added for new functionality
- [x] Manual testing performed

### Test Output

```text
$ uv run --extra dev ruff format --check headroom/proxy/handlers/openai.py
1 file already formatted

$ uv run --extra dev ruff check headroom/proxy/handlers/openai.py tests/test_openai_codex_routing.py tests/test_openai_codex_ws_timings.py tests/test_ws_http_fallback.py
All checks passed!

$ uv run --extra dev python -m pytest -q tests/test_openai_codex_routing.py tests/test_openai_codex_ws_timings.py tests/test_ws_http_fallback.py
37 passed in 0.34s
```

## Real Behavior Proof

- Environment: Local checkout of `fix/codex-store-false-memory-tools`
using `uv run --extra dev`.
- Exact command / steps: Ran the focused formatter, lint, and pytest
commands listed in `Testing`.
- Observed result: Formatting is stable, lint passes, and the focused
OpenAI/Codex routing and fallback tests pass.
- Not tested: Full test suite, `mypy headroom`, and a fresh live ChatGPT
backend probe after the formatter-only follow-up. The original PR
validation recorded that valid ChatGPT subscription backend requests
return `200` with `store=false`, while identical `store=true` or omitted
`store` requests return `400 Store must be set to false`.

## 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
- [ ] 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 or that my
feature works
- [x] New and existing unit tests pass locally with my changes
- [ ] I have updated the CHANGELOG.md if applicable

## Screenshots (if applicable)

N/A.

## Additional Notes

- Post-deploy monitoring terms: `Responses: forced store=false for
ChatGPT auth`, `WS Responses: forced store=false for ChatGPT auth`,
`chatgpt_store_false`, `Memory: forced store=true for Responses memory
tool continuation`, and upstream 400s containing `Store must be set to
false`.
- Expected healthy signals: ChatGPT-auth Responses requests keep
`store=false` and no longer fail with `Store must be set to false`;
API-key memory-tool flows still inject memory tools and can continue via
`previous_response_id`.
- Rollback trigger: any increase in ChatGPT-auth 400s, API-key
memory-tool continuation failures, or missing memory tool injection on
API-key Responses requests.

---------

Co-authored-by: JerrettDavis <mxjerrett@gmail.com>
2026-07-15 20:52:01 +00:00

645 lines
21 KiB
Python

import asyncio
import base64
import json
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import anyio
import pytest
from fastapi import Request
from headroom.proxy.handlers.openai import (
OpenAIHandlerMixin,
_is_allowed_websocket_origin,
_openai_responses_unit_cache_key,
_resolve_codex_routing_headers,
)
def _jwt(payload: dict) -> str:
header = {"alg": "none", "typ": "JWT"}
def encode(part: dict) -> str:
raw = json.dumps(part, separators=(",", ":")).encode("utf-8")
return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=")
return f"{encode(header)}.{encode(payload)}."
def test_resolve_codex_routing_prefers_explicit_header():
headers, is_chatgpt = _resolve_codex_routing_headers(
{
"Authorization": "Bearer sk-test",
"ChatGPT-Account-ID": "acct-explicit",
}
)
assert is_chatgpt is True
assert headers["ChatGPT-Account-ID"] == "acct-explicit"
def test_resolve_codex_routing_derives_account_id_from_oauth_jwt():
token = _jwt(
{
"https://api.openai.com/auth": {
"chatgpt_account_id": "acct-from-jwt",
}
}
)
headers, is_chatgpt = _resolve_codex_routing_headers(
{
"authorization": f"Bearer {token}",
}
)
assert is_chatgpt is True
assert headers["ChatGPT-Account-ID"] == "acct-from-jwt"
def test_resolve_codex_routing_leaves_regular_openai_bearer_tokens_unchanged():
token = _jwt({"aud": ["https://api.openai.com/v1"]})
headers, is_chatgpt = _resolve_codex_routing_headers(
{
"authorization": f"Bearer {token}",
}
)
assert is_chatgpt is False
assert "ChatGPT-Account-ID" not in headers
def test_resolve_codex_routing_returns_none_without_bearer_auth():
headers, is_chatgpt = _resolve_codex_routing_headers({})
assert is_chatgpt is False
assert headers == {}
def test_resolve_codex_routing_ignores_non_jwt_bearer_tokens():
headers, is_chatgpt = _resolve_codex_routing_headers(
{
"authorization": "Bearer not-a-jwt",
}
)
assert is_chatgpt is False
assert headers["authorization"] == "Bearer not-a-jwt"
def test_resolve_codex_routing_ignores_invalid_jwt_payloads():
invalid_payload = base64.urlsafe_b64encode(b"not-json").decode("ascii").rstrip("=")
token = f"test-header.{invalid_payload}.signature"
headers, is_chatgpt = _resolve_codex_routing_headers(
{
"authorization": f"Bearer {token}",
}
)
assert is_chatgpt is False
assert headers["authorization"] == f"Bearer {token}"
def test_openai_responses_unit_cache_key_includes_target_ratio() -> None:
unit = SimpleNamespace(
text="large tool output",
provider="openai",
endpoint="responses",
role="tool",
item_type="function_call_output",
cache_zone="live",
mutable=True,
min_bytes=100,
context=None,
question=None,
bias=None,
metadata={},
)
default_key = _openai_responses_unit_cache_key(unit, model="gpt-5.4")
aggressive_key = _openai_responses_unit_cache_key(
unit,
model="gpt-5.4",
target_ratio=0.10,
)
balanced_key = _openai_responses_unit_cache_key(
unit,
model="gpt-5.4",
target_ratio=0.50,
)
assert aggressive_key != default_key
assert aggressive_key != balanced_key
class _DummyMetrics:
async def record_request(self, **kwargs): # noqa: ANN003
return None
async def record_failed(self, **kwargs): # noqa: ANN003
return None
class _DummyTokenizer:
def count_messages(self, messages):
return len(messages)
class _ResponseStub:
status_code = 200
headers = {"content-type": "application/json", "content-length": "42"}
content = b'{"id":"resp_123","output":[{"type":"message"}]}'
def json(self):
return {"usage": {"input_tokens": 2, "output_tokens": 1}}
class _DummyOpenAIHandler(OpenAIHandlerMixin):
OPENAI_API_URL = "https://api.openai.com"
def __init__(self) -> None:
self.rate_limiter = None
self.metrics = _DummyMetrics()
self.config = SimpleNamespace(
optimize=False,
retry_max_attempts=3,
retry_base_delay_ms=10,
retry_max_delay_ms=50,
connect_timeout_seconds=10,
openai_extra_headers=None,
)
self.usage_reporter = None
self.openai_provider = SimpleNamespace(get_context_limit=lambda model: 128_000)
self.openai_pipeline = SimpleNamespace(apply=MagicMock())
self.anthropic_backend = None
self.cost_tracker = None
self.memory_handler = None
self.traffic_learner = None
# PR-A6 wires session-sticky `OpenAI-Beta` merging into the
# responses HTTP handler — it reads `compute_session_id` to key
# the SessionBetaTracker. The routing tests don't exercise the
# tracker semantics themselves, so a fixed-id stub is enough.
self.session_tracker_store = SimpleNamespace(
compute_session_id=lambda *a, **k: "sess-openai-1",
)
self.captured_request: tuple[str, str, dict, dict] | None = None
self.captured_stream_request: tuple[str, dict, dict] | None = None
async def _next_request_id(self) -> str:
return "req-1"
def _extract_tags(self, headers: dict[str, str]) -> dict[str, str]:
return {}
async def _retry_request(self, method: str, url: str, headers: dict, body: dict, **kwargs):
self.captured_request = (method, url, headers, body)
return _ResponseStub()
async def _run_compression_in_executor(self, fn, *, timeout: float):
# Test stub for HeadroomProxy._run_compression_in_executor.
# The real implementation runs `fn` on a bounded thread pool with
# a wall-clock timeout; tests just need the callable invoked
# synchronously so MagicMock call_count assertions fire.
return fn()
async def _record_request_outcome(self, outcome) -> None:
# Test stub: delegates to the production funnel so wire shape
# matches HeadroomProxy._record_request_outcome.
from headroom.proxy.outcome import emit_request_outcome
await emit_request_outcome(self, outcome)
async def _stream_response(
self,
url: str,
headers: dict,
body: dict,
provider: str,
model: str,
request_id: str,
original_tokens: int,
optimized_tokens: int,
tokens_saved: int,
transforms_applied: list[str],
tags: dict[str, str],
optimization_latency: float,
memory_user_id: str | None = None,
**kwargs,
):
self.captured_stream_request = (url, headers, body)
return SimpleNamespace(
status_code=200,
url=url,
headers=headers,
body=body,
memory_user_id=memory_user_id,
)
class _MemoryToolsOnlyHandler:
def __init__(self) -> None:
self.config = SimpleNamespace(
inject_context=False,
inject_tools=True,
project_root_override="",
)
self.compute_calls = 0
def compute_memory_tool_definitions(self, provider: str) -> list[dict]:
self.compute_calls += 1
assert provider == "openai"
return [
{
"type": "function",
"function": {
"name": "memory_search",
"description": "Search memory.",
"parameters": {"type": "object", "properties": {}},
},
}
]
def has_memory_tool_calls(self, response: dict, provider: str) -> bool:
return False
def _build_request(body: dict, headers: dict[str, str]) -> Request:
payload = json.dumps(body).encode("utf-8")
async def receive():
return {"type": "http.request", "body": payload, "more_body": False}
scope = {
"type": "http",
"asgi": {"version": "3.0"},
"http_version": "1.1",
"method": "POST",
"scheme": "https",
"path": "/v1/responses",
"raw_path": b"/v1/responses",
"query_string": b"",
"headers": [
(key.lower().encode("utf-8"), value.encode("utf-8")) for key, value in headers.items()
],
"client": ("127.0.0.1", 12345),
"server": ("testserver", 443),
}
return Request(scope, receive)
def test_handle_openai_responses_routes_chatgpt_auth_to_backend_api(monkeypatch):
token = _jwt(
{
"https://api.openai.com/auth": {
"chatgpt_account_id": "acct-from-jwt",
}
}
)
request = _build_request(
{"model": "gpt-5.4", "input": "hello"},
{"Authorization": f"Bearer {token}"},
)
handler = _DummyOpenAIHandler()
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert handler.captured_request is not None
method, url, headers, body = handler.captured_request
assert method == "POST"
assert url == "https://chatgpt.com/backend-api/codex/responses"
assert headers["ChatGPT-Account-ID"] == "acct-from-jwt"
assert body["input"] == "hello"
assert body["store"] is False
assert response.status_code == 200
def test_handle_openai_responses_strips_codex_lite_header_upstream(monkeypatch):
# OpenAI rejects newer Codex models when the client-only lite header leaks
# upstream. The HTTP POST path must drop it like the WS handler does, while
# leaving adjacent headers intact.
token = _jwt(
{
"https://api.openai.com/auth": {
"chatgpt_account_id": "acct-from-jwt",
}
}
)
request = _build_request(
{"model": "gpt-5.4", "input": "hello"},
{
"Authorization": f"Bearer {token}",
"X-OpenAI-Internal-Codex-Responses-Lite": "true",
"X-OpenAI-Debug": "keep-me",
},
)
handler = _DummyOpenAIHandler()
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
_method, _url, headers, _body = handler.captured_request
lowered = {k.lower(): v for k, v in headers.items()}
assert "x-openai-internal-codex-responses-lite" not in lowered
assert lowered.get("x-openai-debug") == "keep-me"
def test_handle_openai_responses_chatgpt_auth_skips_memory_tools(monkeypatch):
token = _jwt(
{
"https://api.openai.com/auth": {
"chatgpt_account_id": "acct-from-jwt",
}
}
)
request = _build_request(
{"model": "gpt-5.4", "input": "hello", "store": True},
{"Authorization": f"Bearer {token}", "x-headroom-user-id": "user-1"},
)
handler = _DummyOpenAIHandler()
memory_handler = _MemoryToolsOnlyHandler()
handler.memory_handler = memory_handler
handler.session_tracker_store = SimpleNamespace(
compute_session_id=lambda *a, **k: "sess-chatgpt-no-memory-tools",
)
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
_, url, _, body = handler.captured_request
assert url == "https://chatgpt.com/backend-api/codex/responses"
assert body["store"] is False
assert "tools" not in body
assert memory_handler.compute_calls == 0
def test_handle_openai_responses_chatgpt_codex_timeout_fails_open(monkeypatch):
token = _jwt(
{
"https://api.openai.com/auth": {
"chatgpt_account_id": "acct-from-jwt",
}
}
)
request = _build_request(
{"model": "gpt-5.4", "input": "large context"},
{"Authorization": f"Bearer {token}"},
)
handler = _DummyOpenAIHandler()
handler.config.optimize = True
async def timeout_compression(*args, **kwargs): # noqa: ANN002, ANN003
raise asyncio.TimeoutError()
handler._compress_openai_responses_payload_in_executor = timeout_compression
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
method, url, headers, body = handler.captured_request
assert method == "POST"
assert url == "https://chatgpt.com/backend-api/codex/responses"
assert body["input"] == "large context"
assert body["store"] is False
def test_handle_openai_responses_api_auth_store_false_skips_memory_tools(monkeypatch):
request = _build_request(
{"model": "gpt-4o-mini", "input": "hello", "store": False},
{"Authorization": "Bearer sk-test", "x-headroom-user-id": "user-1"},
)
handler = _DummyOpenAIHandler()
memory_handler = _MemoryToolsOnlyHandler()
handler.memory_handler = memory_handler
handler.session_tracker_store = SimpleNamespace(
compute_session_id=lambda *a, **k: "sess-api-memory-tools",
)
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
_, url, _, body = handler.captured_request
assert url == "https://api.openai.com/v1/responses"
assert body["store"] is False
assert "tools" not in body
assert memory_handler.compute_calls == 1
def test_handle_openai_responses_routes_api_key_auth_direct_to_openai(monkeypatch):
request = _build_request(
{"model": "gpt-4o-mini", "input": "hello"},
{"Authorization": "Bearer sk-test"},
)
handler = _DummyOpenAIHandler()
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert handler.captured_request is not None
method, url, headers, body = handler.captured_request
assert method == "POST"
assert url == "https://api.openai.com/v1/responses"
assert headers.get("ChatGPT-Account-ID") is None
assert body["input"] == "hello"
assert response.status_code == 200
def test_handle_openai_responses_stream_skips_python_compression(monkeypatch):
"""PR-C5: Python no longer compresses /v1/responses (Rust handles it
natively). The streaming forward path must still fire — only the
Python compression dispatch is retired."""
request = _build_request(
{
"model": "gpt-5.4",
"stream": True,
"instructions": "Keep it short",
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "hello"}],
}
],
},
{"Authorization": "Bearer sk-test"},
)
handler = _DummyOpenAIHandler()
handler.config.optimize = True
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_stream_request is not None
assert handler.openai_pipeline.apply.call_count == 0
assert handler.captured_stream_request[2]["stream"] is True
def test_handle_openai_responses_memory_timeout_fails_open(monkeypatch):
class _SlowMemoryHandler:
def __init__(self):
self.config = SimpleNamespace(inject_context=True, inject_tools=False)
async def search_and_format_context(self, memory_user_id, messages, **_kwargs):
return "should not be used"
def has_memory_tool_calls(self, response, provider):
return False
async def _timeout_wait_for(awaitable, timeout):
close = getattr(awaitable, "close", None)
if callable(close):
close()
raise TimeoutError
request = _build_request(
{"model": "gpt-5.4", "input": "hello"},
{"Authorization": "Bearer sk-test", "x-headroom-user-id": "user-1"},
)
handler = _DummyOpenAIHandler()
handler.memory_handler = _SlowMemoryHandler()
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
monkeypatch.setattr("headroom.proxy.handlers.openai.asyncio.wait_for", _timeout_wait_for)
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
_, _, _, body = handler.captured_request
assert body.get("instructions") is None
def test_codex_responses_timeout_fails_open_in_standalone_proxy(monkeypatch):
"""Codex users running only the proxy still get fail-open on timeout."""
request = _build_request(
{
"model": "gpt-5.4",
"input": [
{
"type": "function_call_output",
"call_id": "call-1",
"output": "large tool output",
}
],
},
{"Authorization": "Bearer sk-test", "x-client": "codex"},
)
handler = _DummyOpenAIHandler()
handler.config.optimize = True
monkeypatch.setattr("headroom.tokenizers.get_tokenizer", lambda model: _DummyTokenizer())
monkeypatch.setattr(
handler,
"_compress_openai_responses_payload",
lambda *args, **kwargs: (_ for _ in ()).throw(TimeoutError()),
)
response = anyio.run(handler.handle_openai_responses, request)
assert response.status_code == 200
assert handler.captured_request is not None
_, url, _, body = handler.captured_request
assert url == "https://api.openai.com/v1/responses"
assert body["input"][0]["output"] == "large tool output"
class _DummyWebSocket:
def __init__(self, headers: dict[str, str]):
self.headers = headers
self.accepted_subprotocol = None
self.closed = False
self.close_code = None
self.close_reason = None
async def accept(self, subprotocol=None, headers=None):
self.accepted_subprotocol = subprotocol
async def close(self, code=1000, reason=None):
self.closed = True
self.close_code = code
self.close_reason = reason
def test_websocket_origin_policy_allows_native_clients_without_origin(monkeypatch):
monkeypatch.delenv("HEADROOM_WS_ORIGINS", raising=False)
monkeypatch.delenv("HEADROOM_CORS_ORIGINS", raising=False)
assert _is_allowed_websocket_origin({"authorization": "Bearer token"}) is True
def test_websocket_origin_policy_allows_loopback_origins_by_default(monkeypatch):
monkeypatch.delenv("HEADROOM_WS_ORIGINS", raising=False)
monkeypatch.delenv("HEADROOM_CORS_ORIGINS", raising=False)
assert _is_allowed_websocket_origin({"origin": "http://localhost:3000"}) is True
assert _is_allowed_websocket_origin({"origin": "https://127.0.0.1:8787"}) is True
def test_websocket_origin_policy_requires_config_for_remote_origins(monkeypatch):
monkeypatch.delenv("HEADROOM_WS_ORIGINS", raising=False)
monkeypatch.delenv("HEADROOM_CORS_ORIGINS", raising=False)
assert _is_allowed_websocket_origin({"origin": "https://remote.example"}) is False
assert _is_allowed_websocket_origin({"origin": "http://"}) is False
def test_websocket_origin_policy_can_be_pinned_with_env(monkeypatch):
monkeypatch.setenv("HEADROOM_WS_ORIGINS", "https://dash.example.com")
monkeypatch.delenv("HEADROOM_CORS_ORIGINS", raising=False)
assert _is_allowed_websocket_origin({"origin": "https://dash.example.com"}) is True
assert _is_allowed_websocket_origin({"origin": "http://localhost:3000"}) is False
def test_handle_openai_responses_ws_resolves_codex_routing_headers():
class SentinelError(RuntimeError):
pass
handler = _DummyOpenAIHandler()
websocket = _DummyWebSocket({"authorization": "Bearer token"})
with patch.dict(sys.modules, {"websockets": MagicMock()}):
with patch(
"headroom.proxy.handlers.openai._resolve_codex_routing_headers",
side_effect=SentinelError("resolved"),
):
with pytest.raises(SentinelError, match="resolved"):
anyio.run(handler.handle_openai_responses_ws, websocket)
def test_handle_openai_responses_ws_closes_unconfigured_origin(monkeypatch):
handler = _DummyOpenAIHandler()
websocket = _DummyWebSocket({"origin": "https://remote.example"})
monkeypatch.delenv("HEADROOM_WS_ORIGINS", raising=False)
monkeypatch.delenv("HEADROOM_CORS_ORIGINS", raising=False)
with patch.dict(sys.modules, {"websockets": MagicMock()}):
with patch(
"headroom.proxy.handlers.openai._resolve_codex_routing_headers",
side_effect=AssertionError("routing should not run"),
):
anyio.run(handler.handle_openai_responses_ws, websocket)
assert websocket.closed is True
assert websocket.close_code == 1008
assert websocket.close_reason == "origin not allowed"
assert websocket.accepted_subprotocol is None