fix(proxy): preserve Responses passthrough bytes

This commit is contained in:
Vinay Gupta 2026-06-30 09:33:22 -04:00
parent 1c0e15243e
commit 66507d98d4
3 changed files with 184 additions and 4 deletions

View file

@ -2861,7 +2861,8 @@ class OpenAIHandlerMixin:
from headroom.proxy.helpers import (
MAX_REQUEST_BODY_SIZE,
_read_request_json,
BodyMutationTracker,
read_request_json_with_bytes,
)
from headroom.tokenizers import get_tokenizer
from headroom.utils import extract_user_query
@ -2893,7 +2894,7 @@ class OpenAIHandlerMixin:
# Parse request
try:
body = await _read_request_json(request)
body, original_body_bytes = await read_request_json_with_bytes(request)
except (json.JSONDecodeError, ValueError) as e:
return JSONResponse(
status_code=400,
@ -2908,6 +2909,7 @@ class OpenAIHandlerMixin:
model = body.get("model", "unknown")
stream = body.get("stream", False)
body_mutation_tracker = BodyMutationTracker()
_bypass = self._headroom_bypass_enabled(request.headers)
if _bypass:
logger.info(
@ -2948,6 +2950,10 @@ class OpenAIHandlerMixin:
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
# The parsed request body has already been content-decoded. Remove
# entity headers that described the client-to-proxy wire body.
headers.pop("content-encoding", None)
headers.pop("transfer-encoding", None)
# Strip accept-encoding so httpx negotiates its own encoding.
# Cloudflare Workers forward "br, zstd" which OpenAI may honor;
# if httpx lacks brotli support the response body is undecipherable → 502.
@ -3134,6 +3140,7 @@ class OpenAIHandlerMixin:
if current_input
else memory_context
)
body_mutation_tracker.mark_mutated("responses_memory_context")
log_memory_injection(
request_id=request_id,
session_id=None,
@ -3147,6 +3154,7 @@ class OpenAIHandlerMixin:
)
if bytes_appended > 0:
body["input"] = new_input
body_mutation_tracker.mark_mutated("responses_memory_context")
log_memory_injection(
request_id=request_id,
session_id=None,
@ -3211,12 +3219,14 @@ class OpenAIHandlerMixin:
)
if mem_tools_injected:
body["tools"] = resp_tools
body_mutation_tracker.mark_mutated("responses_memory_tools")
logger.info(f"[{request_id}] Memory: Injected memory tools (openai/responses)")
if _ensure_responses_store_for_memory_tools(
body,
memory_tools_injected=True,
):
body_mutation_tracker.mark_mutated("responses_memory_store")
logger.info(
f"[{request_id}] Memory: forced store=true for Responses memory tool continuation"
)
@ -3272,6 +3282,7 @@ class OpenAIHandlerMixin:
)
attempted_input_tokens = int(_attempted_tokens)
if _modified:
body_mutation_tracker.mark_mutated("responses_compression")
tokens_saved = int(_tokens_saved)
optimized_tokens = max(0, original_tokens - tokens_saved)
transforms_applied = [*_transforms, *list(transforms_applied)]
@ -3397,11 +3408,25 @@ class OpenAIHandlerMixin:
optimization_latency,
memory_user_id=memory_user_id,
memory_request_ctx=memory_request_ctx,
original_body_bytes=original_body_bytes,
body_mutated=body_mutation_tracker.mutated,
mutation_reasons=body_mutation_tracker.reasons,
waste_signals=waste_signals_dict,
)
else:
headers = await apply_copilot_api_auth(headers, url=url)
response = await self._retry_request("POST", url, headers, body)
response = await self._retry_request(
"POST",
url,
headers,
body,
original_body_bytes=original_body_bytes,
body_mutated=body_mutation_tracker.mutated,
mutation_reasons=body_mutation_tracker.reasons,
request_id=request_id,
forwarder_name="openai_responses",
path_for_log=url,
)
_response_body_for_debug: Any = None
_response_raw_for_debug: str | None = None
try:

View file

@ -192,7 +192,7 @@ class _DummyOpenAIHandler(OpenAIHandlerMixin):
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):
async def _retry_request(self, method: str, url: str, headers: dict, body: dict, **kwargs):
self.captured_request = (method, url, headers, body)
return _ResponseStub()

View file

@ -18,8 +18,10 @@ rollback (operator opt-in, not a fallback).
from __future__ import annotations
import gzip
import hashlib
import json
import logging
from types import SimpleNamespace
from unittest.mock import AsyncMock
@ -340,6 +342,92 @@ def _make_no_optimize_app() -> tuple[TestClient, _CapturingTransport]:
return _make_anthropic_app(optimize=False)
def _openai_responses_body_bytes(*, stream: bool) -> bytes:
payload = {
"model": "gpt-5.5",
"input": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_text",
"text": "hello 🔥 with spaces preserved",
}
],
}
],
"stream": stream,
}
return json.dumps(payload, ensure_ascii=False, indent=2).encode("utf-8")
def _openai_responses_codex_headers(content_encoding: str) -> dict[str, str]:
return {
"authorization": "Bearer test-token",
"chatgpt-account-id": "acct_test",
"originator": "Codex Desktop",
"content-type": "application/json",
"content-encoding": content_encoding,
"accept": "text/event-stream",
}
def _start_proxy_log_capture() -> tuple[
logging.Logger,
logging.Handler,
int,
list[logging.LogRecord],
]:
proxy_logger = logging.getLogger("headroom.proxy")
records: list[logging.LogRecord] = []
class _ListHandler(logging.Handler):
def emit(self, record: logging.LogRecord) -> None:
records.append(record)
handler = _ListHandler(level=logging.INFO)
prev_level = proxy_logger.level
proxy_logger.addHandler(handler)
proxy_logger.setLevel(logging.INFO)
return proxy_logger, handler, prev_level, records
def _stop_proxy_log_capture(
proxy_logger: logging.Logger,
handler: logging.Handler,
prev_level: int,
) -> None:
proxy_logger.removeHandler(handler)
proxy_logger.setLevel(prev_level)
def _assert_openai_responses_encoded_passthrough(
transport: _CapturingTransport,
decoded_body: bytes,
) -> None:
assert transport.captured_body == decoded_body
assert transport.captured_headers is not None
captured_headers = {key.lower(): value for key, value in transport.captured_headers.items()}
assert "content-encoding" not in captured_headers
assert captured_headers.get("content-length") == str(len(decoded_body))
def _assert_outbound_passthrough_log(
records: list[logging.LogRecord],
*,
forwarder: str,
) -> None:
messages = [record.getMessage() for record in records]
assert any(
"event=outbound_request" in message
and f"forwarder={forwarder}" in message
and "body_mutated=false" in message
and "source=passthrough" in message
for message in messages
), messages
def test_passthrough_no_mutation_byte_equal_sha256() -> None:
"""No transform → upstream SHA-256 equals client-sent SHA-256."""
client, transport = _make_no_optimize_app()
@ -904,6 +992,73 @@ def test_streaming_forwarder_byte_faithful() -> None:
)
def test_openai_responses_gzip_nonstream_passthrough_strips_content_encoding() -> None:
client, transport = _make_no_optimize_app()
decoded_body = _openai_responses_body_bytes(stream=False)
encoded_body = gzip.compress(decoded_body)
proxy_logger, handler, prev_level, records = _start_proxy_log_capture()
try:
response = client.post(
"/v1/responses",
headers=_openai_responses_codex_headers("gzip"),
content=encoded_body,
)
finally:
_stop_proxy_log_capture(proxy_logger, handler, prev_level)
assert response.status_code == 200, response.text
_assert_openai_responses_encoded_passthrough(transport, decoded_body)
_assert_outbound_passthrough_log(records, forwarder="openai_responses")
def test_openai_responses_gzip_stream_passthrough_strips_content_encoding() -> None:
client, transport = _make_no_optimize_app()
decoded_body = _openai_responses_body_bytes(stream=True)
encoded_body = gzip.compress(decoded_body)
proxy_logger, handler, prev_level, records = _start_proxy_log_capture()
try:
with client.stream(
"POST",
"/v1/responses",
headers=_openai_responses_codex_headers("gzip"),
content=encoded_body,
) as response:
assert response.status_code == 200
for _ in response.iter_bytes():
pass
finally:
_stop_proxy_log_capture(proxy_logger, handler, prev_level)
_assert_openai_responses_encoded_passthrough(transport, decoded_body)
_assert_outbound_passthrough_log(records, forwarder="streaming")
def test_openai_responses_codex_desktop_zstd_stream_passthrough_strips_content_encoding() -> None:
zstandard = pytest.importorskip("zstandard")
client, transport = _make_no_optimize_app()
decoded_body = _openai_responses_body_bytes(stream=True)
encoded_body = zstandard.ZstdCompressor().compress(decoded_body)
proxy_logger, handler, prev_level, records = _start_proxy_log_capture()
try:
with client.stream(
"POST",
"/v1/responses",
headers=_openai_responses_codex_headers("zstd"),
content=encoded_body,
) as response:
assert response.status_code == 200
for _ in response.iter_bytes():
pass
finally:
_stop_proxy_log_capture(proxy_logger, handler, prev_level)
_assert_openai_responses_encoded_passthrough(transport, decoded_body)
_assert_outbound_passthrough_log(records, forwarder="streaming")
# ---------------------------------------------------------------------------
# Batch forwarder byte-faithfulness (passthrough variant)
# ---------------------------------------------------------------------------