mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
fix(proxy): preserve Responses passthrough bytes
This commit is contained in:
parent
1c0e15243e
commit
66507d98d4
3 changed files with 184 additions and 4 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue