headroom/tests/test_proxy_handlers_batch.py
Tejas Chopra f9807fd69e
feat(proxy): let extensions report cost savings and their own latency (#3051)
## What

Two changes that let a proxy extension report **what it saved** and
**what it cost**, so both show up under `/stats`, the dashboard, and
Prometheus.

`record_scope_savings` already existed and already accepted `usd` — the
one channel in the proxy that can express savings *without* tokens. Two
things stopped it working end to end.

### 1. Savings were silently dropped on Gemini traffic (bug)

`bind_scope` shares one attribution ledger between ASGI middleware and
the request handler. Anthropic and OpenAI call it; **Gemini never did**,
so anything an extension recorded into the request scope was discarded
for Gemini traffic only — silently, because an empty ledger and an
unbound one are indistinguishable at the outcome funnel. Now bound at
all four Gemini tag sites.

### 2. An extension's own latency was invisible (gap)

`overhead_ms` is measured *inside* the handler, and an ASGI extension
**wraps** that handler — so every millisecond it spends reaches the
client while every timing surface stays flat. An extension that halves
the bill and adds 200 ms per request is a trade the operator has to see
both halves of, and only one half was reaching the dashboard.

`record_scope_timing(scope, stage, ms)` is the symmetric counterpart to
`record_scope_savings`, carried on the same bound ledger and merged into
`RequestOutcome.pipeline_timing` at the outcome funnel — one place, so
every provider picks it up at once.

## API surface

```python
from headroom.proxy.savings_attribution import record_scope_savings, record_scope_timing

record_scope_savings(scope, "my_extension", tokens=0, usd=0.004)   # money without tokens
record_scope_timing(scope, "my_extension", elapsed_ms)
```

Both take the ASGI `scope`, because middleware has no other way in.
Documented in `extensions.py` — the module extension authors actually
read, and the stability contract for this interface.

- Savings → `/stats` `savings.by_source`, dashboard card,
`headroom_savings_attributed_usd_total{source=...}`
- Timing → `/stats` `pipeline_timing`, dashboard Performance panel,
`headroom_transform_timing_ms_*`

**Attribution only.** These rows explain the headline total; they are
never added to it.

## Changes to existing behavior

- `public_tags` now strips `_headroom_stage_timing` as well as
`_headroom_savings_attribution`. Both ride on `tags` because that is the
one dict reaching the outcome funnel from every handler, and a list and
a dict must not land in a string-keyed label store.
- `pipeline_timing` passed to `metrics.record_request` is merged rather
than passed through **only when an extension contributed timings**; with
no extension the handler's own dict is passed through unchanged
(asserted by identity in the tests).
- Stage names are extension-supplied, so they are capped at 16 and
namespaced `ext:` — `deep_copy` reported by a plugin must never
accumulate into the same series as `deep_copy` measured by the pipeline.
A handler's own timing wins a collision (unreachable while the prefix
stands; the safe way round if it ever goes).

## Failure modes

Both calls are bounded (32 sources, 16 stages), never raise, and never
change a response — telemetry from a plugin must not be able to break
the request it is describing. Non-positive and non-numeric durations are
ignored: a zero is a clock artifact, not an observation, and averaging
it in would drag the mean down exactly where the stage is cheapest to
skip. `timings_from_tags` tolerates junk on the tag.

## Test-double fix

Three Gemini test fakes (`FakeRequest`, `_FakeRequest`,
`_VertexGeminiImageRequest`) had no `.scope`, which every real Starlette
`Request` has. They now do. This is a double that had drifted from the
type it stands in for; the alternative was weakening the handler to
tolerate a request shape that cannot occur in production.

---

## Real behavior proof

**Setup:** macOS 15.4 (darwin 25.4.0), Python 3.12.13, this branch at
`c814b950`, real `create_app` proxy with `respx`-mocked Anthropic
upstream, a demo ASGI extension added via `app.add_middleware`.

**The extension** — written as a third party would, reporting `tokens=0`
because it re-routed `claude-opus-5` → `claude-haiku-4-5`: same tokens,
cheaper model. That is precisely the case no existing Headroom savings
channel can express, since all of them compute `saved = before - after`.

```python
class DemoRouter:
    def __init__(self, app): self.app = app
    async def __call__(self, scope, receive, send):
        if scope.get("type") != "http":
            return await self.app(scope, receive, send)
        started = time.perf_counter()
        record_scope_savings(scope, "routemegood", tokens=0, usd=0.173)
        record_scope_timing(scope, "routemegood", (time.perf_counter() - started) * 1000)
        await self.app(scope, receive, send)
```

**Ran:** three POSTs to `/v1/messages`, then `GET /stats` and `GET
/metrics`.

**Observed:**

```
upstream call -> 200
upstream call -> 200
upstream call -> 200

=== /stats  savings.by_source  (what the dashboard renders) ===
[
  {
    "source": "routemegood",
    "realized": true,
    "events": 3,
    "tokens": 0,
    "usd": 0.519
  }
]

=== /stats  pipeline_timing  (dashboard Performance panel) ===
{
  "ext:routemegood": {
    "average_ms": 0.01,
    "max_ms": 0.02,
    "count": 3
  }
}

=== /metrics ===
# HELP headroom_savings_attributed_tokens_total Tokens attributed to a savings source
# TYPE headroom_savings_attributed_tokens_total counter
headroom_savings_attributed_tokens_total{realized="true",source="routemegood"} 0
# HELP headroom_savings_attributed_usd_total Cost savings attributed to a source; may be negative
# TYPE headroom_savings_attributed_usd_total gauge
headroom_savings_attributed_usd_total{realized="true",source="routemegood"} 0.519
headroom_transform_timing_ms_sum{transform="ext:routemegood"} 0.03
```

`$0.519 = 3 × $0.173` — three requests, correctly accumulated, with
`tokens: 0` throughout.

**Also have (not a substitute for the above):** 22 new unit tests in
`tests/test_extension_attribution.py`, including four that drive the
real `_record_request_outcome` funnel via the same descriptor-binding
harness `test_request_outcome.py` uses.

Full suite on this branch: **10,989 passed, 578 skipped**. Three
failures —
`test_graceful_shutdown.py::test_run_server_installs_cancelled_error_filter`
(full-suite ordering; passes in isolation),
`test_learn/test_integration.py::TestCodexIntegration::test_full_pipeline`,
and `test_release_workflows.py::test_no_native_tls_in_wheel_build_tree`
(needs `cargo`) — **reproduce identically on clean `main`** (`2f4d001c`,
10,967 passed, same 3 failed). Verified by stashing this branch and
re-running the full suite on main in the same tree.

**What I did not test:** a live provider (upstream is `respx`-mocked);
the Gemini `bind_scope` fix against real Google traffic (covered by the
existing 114 Gemini tests, which all pass); the dashboard rendered in a
browser — I verified the JSON shape its templates bind to
(`stats.savings?.by_source`, `stats.pipeline_timing`) rather than the
pixels.

---

🤖 Generated with [Claude Code](https://claude.com/claude-code)

---------

Co-authored-by: Tejas Chopra <tejas@Tejass-MacBook-Pro.local>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-08-16 10:25:47 -07:00

1752 lines
63 KiB
Python

from __future__ import annotations
import json
import sys
from types import SimpleNamespace
import pytest
from headroom.cache.compression_store import (
CompressionEntry,
get_compression_store,
reset_compression_store,
)
from headroom.ccr import response_handler as response_handler_module
from headroom.proxy.handlers import batch as batch_module
from headroom.proxy.handlers import gemini as gemini_module
from headroom.proxy.handlers.gemini import GeminiHandlerMixin
class FakeResponse:
def __init__(
self,
*,
status_code: int = 200,
content: bytes = b"{}",
headers: dict[str, str] | None = None,
text: str | None = None,
json_data=None, # noqa: ANN001
) -> None:
self.status_code = status_code
self.content = content
self.headers = headers or {}
self.text = text if text is not None else content.decode("utf-8", errors="ignore")
self._json_data = json_data
def json(self): # noqa: ANN201
if self._json_data is not None:
return self._json_data
return json.loads(self.text)
class FakeHttpClient:
def __init__(self) -> None:
self.posts: list[dict[str, object]] = []
self.gets: list[dict[str, object]] = []
self.requests: list[dict[str, object]] = []
self.post_response = FakeResponse()
self.get_response = FakeResponse()
self.raise_post: Exception | None = None
self.raise_get: Exception | None = None
async def post(self, url: str, **kwargs): # noqa: ANN003, ANN201
self.posts.append({"url": url, **kwargs})
if self.raise_post is not None:
raise self.raise_post
return self.post_response
async def get(self, url: str, **kwargs): # noqa: ANN003, ANN201
self.gets.append({"url": url, **kwargs})
if self.raise_get is not None:
raise self.raise_get
return self.get_response
async def request(self, method: str, url: str, **kwargs): # noqa: ANN003, ANN201
self.requests.append({"method": method, "url": url, **kwargs})
if self.raise_get is not None:
raise self.raise_get
return self.get_response
class FakeMetrics:
def __init__(self) -> None:
self.record_calls: list[dict[str, object]] = []
self.failed_calls: list[dict[str, object]] = []
async def record_request(self, **kwargs) -> None: # noqa: ANN003
self.record_calls.append(kwargs)
async def record_failed(self, **kwargs) -> None: # noqa: ANN003
self.failed_calls.append(kwargs)
class DummyBatchHandler(batch_module.BatchHandlerMixin, GeminiHandlerMixin):
# GeminiHandlerMixin supplies the real _rebuild_gemini_contents (and the
# other content helpers); the two converter methods below intentionally
# override the mixin's for the stub-based tests.
OPENAI_API_URL = "https://openai.example"
GEMINI_API_URL = "https://gemini.example"
def __init__(self) -> None:
self.http_client = FakeHttpClient()
self.metrics = FakeMetrics()
self.config = SimpleNamespace(
optimize=False,
ccr_inject_tool=False,
ccr_inject_system_instructions=False,
)
self.openai_provider = SimpleNamespace(get_context_limit=lambda model: 8192)
self.openai_pipeline = SimpleNamespace(apply=lambda **kwargs: None)
self._request_counter = 0
self._retry_response = FakeResponse()
async def _next_request_id(self) -> str:
self._request_counter += 1
return f"req-{self._request_counter}"
async def _record_request_outcome(self, outcome) -> None: # noqa: ANN001
# Mirror of HeadroomProxy._record_request_outcome for the batch
# mixin tests. Delegates to the free funnel so the wire shape
# matches production.
from headroom.proxy.outcome import emit_request_outcome
await emit_request_outcome(self, outcome)
def _extract_tags(self, headers: dict) -> dict[str, str]:
# Mirror of HeadroomProxy._extract_tags. Handlers now call this
# at entry to capture x-headroom-* slicing tags into the outcome.
return {
k.lower().replace("x-headroom-", ""): v
for k, v in headers.items()
if k.lower().startswith("x-headroom-")
}
async def handle_passthrough(self, request, base_url): # noqa: ANN001, ANN201
return {"request": request, "base_url": base_url}
async def _run_compression_in_executor(self, fn, *, timeout): # noqa: ANN001, ANN201
# Mirror of HeadroomProxy._run_compression_in_executor: batch handlers
# offload pipeline.apply() off the event loop (#1701). Inline is fine
# for tests — only the call contract matters here.
return fn()
async def _retry_request(self, method, url, headers, body, **kwargs): # noqa: ANN001, ANN201
return self._retry_response
def _gemini_contents_to_messages(self, contents, system_instruction): # noqa: ANN001, ANN201
messages = [{"role": "user", "content": part["parts"][0]["text"]} for part in contents]
return messages, []
def _messages_to_gemini_contents(self, messages): # noqa: ANN001, ANN201
return ([{"parts": [{"text": message["content"]}]} for message in messages], None)
class FakeRequest:
def __init__(
self,
body: bytes | str,
*,
headers: dict[str, str] | None = None,
method: str = "POST",
path: str = "/v1/batches",
query: str = "",
) -> None:
self._body = body.encode("utf-8") if isinstance(body, str) else body
self.headers = headers or {}
self.method = method
self.url = SimpleNamespace(path=path, query=query)
self.query_params = {}
# Every real Starlette Request has one, and handlers now share a
# per-request attribution ledger through it (savings_attribution).
self.scope: dict = {"type": "http", "method": method}
async def body(self) -> bytes:
return self._body
class NativeGeminiHandler(DummyBatchHandler):
def __init__(self, responses: list[FakeResponse]) -> None:
super().__init__()
self.config.optimize = True
self.config.ccr_inject_tool = True
self.config.ccr_inject_system_instructions = False
self.memory_handler = None
self.rate_limiter = None
self.usage_reporter = None
self.responses = iter(responses)
self.sent_bodies: list[dict] = []
from headroom.ccr.response_handler import CCRResponseHandler
self.ccr_response_handler = CCRResponseHandler()
self.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: SimpleNamespace(
messages=[
{
"role": "user",
"content": "compressed [100 items compressed to 1. Retrieve more: hash=aaaaaaaaaaaaaaaaaaaaaaaa]",
}
],
timing={},
tokens_before=10,
tokens_after=5,
transforms_applied=[],
waste_signals=SimpleNamespace(to_dict=lambda: {}),
)
)
def _gemini_contents_to_messages(
self, contents, system_instruction=None, *, include_function_responses=False
): # noqa: ANN001, ANN201
return GeminiHandlerMixin._gemini_contents_to_messages(
self,
contents,
system_instruction,
include_function_responses=include_function_responses,
)
def _messages_to_gemini_contents(self, messages): # noqa: ANN001, ANN201
return GeminiHandlerMixin._messages_to_gemini_contents(self, messages)
async def _retry_request(self, method, url, headers, body, **kwargs): # noqa: ANN001, ANN201
self.sent_bodies.append(body)
return next(self.responses)
async def _run_compression_in_executor(self, fn, *, timeout): # noqa: ANN001, ANN201
return fn()
def install_native_gemini_compression(monkeypatch: pytest.MonkeyPatch) -> None:
class Decision:
should_compress = True
passthrough_reason = ""
def apply_to_tags(self, tags) -> None: # noqa: ANN001
return None
monkeypatch.setattr(gemini_module.CompressionDecision, "decide", lambda **kwargs: Decision())
def native_gemini_request(tools=None) -> dict: # noqa: ANN001
return {
"contents": [{"role": "user", "parts": [{"text": "compressed input"}]}],
"generationConfig": {"temperature": 0.2},
**({"tools": tools} if tools is not None else {}),
}
def native_ccr_response() -> FakeResponse:
return FakeResponse(
json_data={
"candidates": [
{
"content": {
"role": "model",
"parts": [
{
"functionCall": {
"name": "headroom_retrieve",
"id": "call-1",
"args": {"hash": "aaaaaaaaaaaaaaaaaaaaaaaa"},
}
}
],
}
}
],
"usageMetadata": {"promptTokenCount": 5},
}
)
@pytest.mark.asyncio
async def test_gemini_native_ccr_continuation(monkeypatch: pytest.MonkeyPatch) -> None:
install_native_gemini_compression(monkeypatch)
from headroom.ccr.response_handler import CCRToolResult
final = FakeResponse(
json_data={
"candidates": [{"content": {"role": "model", "parts": [{"text": "final answer"}]}}]
}
)
handler = NativeGeminiHandler([native_ccr_response(), final])
handler.ccr_response_handler._execute_retrieval = lambda call: CCRToolResult(
call.tool_call_id,
json.dumps({"hash": call.hash_key, "original_content": [{"type": "code"}]}),
True,
1,
"headroom_retrieve",
)
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json", "x-goog-api-key": "secret"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 200
assert (
json.loads(response.body)["candidates"][0]["content"]["parts"][0]["text"] == "final answer"
), response.body
assert len(handler.sent_bodies) == 2
continuation = handler.sent_bodies[1]["contents"]
assert continuation[-2]["role"] == "model"
assert continuation[-2]["parts"][0]["functionCall"]["name"] == "headroom_retrieve"
assert continuation[-1]["role"] == "user"
assert continuation[-1]["parts"][0]["functionResponse"]["name"] == "headroom_retrieve"
assert continuation[-1]["parts"][0]["functionResponse"]["id"] == "call-1"
@pytest.mark.asyncio
async def test_gemini_native_ccr_tools(monkeypatch: pytest.MonkeyPatch) -> None:
install_native_gemini_compression(monkeypatch)
# verify_ownership() (issue #2836) requires the marker's hash to be a
# real store entry; NativeGeminiHandler's mocked pipeline hand-types
# "hash=aaaa...aaaa" rather than compressing through the real store.
reset_compression_store()
get_compression_store().store(
original="original content",
compressed="compressed [100 items compressed to 1]",
explicit_hash="aaaaaaaaaaaaaaaaaaaaaaaa",
)
handler = NativeGeminiHandler(
[FakeResponse(json_data={"candidates": [{"content": {"parts": [{"text": "answer"}]}}]})]
)
tools = [
{"functionDeclarations": [{"name": "client_tool"}]},
{"functionDeclarations": [{"name": "second_tool"}]},
{"googleSearch": {}},
{"codeExecution": {}},
]
await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request(tools)),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
forwarded_tools = handler.sent_bodies[0]["tools"]
assert forwarded_tools[2:] == tools[2:]
declarations = forwarded_tools[0]["functionDeclarations"]
assert {item["name"] for item in declarations} == {"client_tool", "headroom_retrieve"}
assert forwarded_tools[1]["functionDeclarations"] == [{"name": "second_tool"}]
reset_compression_store()
@pytest.mark.asyncio
async def test_gemini_native_ccr_does_not_duplicate_existing_declaration(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
tools = [
{"functionDeclarations": [{"name": "client_tool"}]},
{"functionDeclarations": [{"name": "headroom_retrieve"}]},
]
handler = NativeGeminiHandler(
[FakeResponse(json_data={"candidates": [{"content": {"parts": [{"text": "answer"}]}}]})]
)
await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request(tools)),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
names = [
declaration["name"]
for tool in handler.sent_bodies[0]["tools"]
for declaration in tool.get("functionDeclarations", [])
]
assert names.count("headroom_retrieve") == 1
@pytest.mark.asyncio
async def test_gemini_native_ccr_does_not_inject_into_streaming_request(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
handler = NativeGeminiHandler([FakeResponse()])
captured: dict[str, object] = {}
async def fake_stream(*args, **kwargs): # noqa: ANN002, ANN003, ANN202
captured["body"] = args[2]
return FakeResponse()
monkeypatch.setattr(handler, "_stream_response", fake_stream, raising=False)
tools = [{"functionDeclarations": [{"name": "client_tool"}]}]
await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request(tools)),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:streamGenerateContent",
),
"gemini-2.5-flash",
)
streamed_tools = captured["body"]["tools"] # type: ignore[index]
names = [
declaration["name"]
for tool in streamed_tools
for declaration in tool.get("functionDeclarations", [])
]
assert names == ["client_tool"]
@pytest.mark.asyncio
async def test_gemini_native_ccr_mixed(monkeypatch: pytest.MonkeyPatch) -> None:
install_native_gemini_compression(monkeypatch)
response_json = {
"candidates": [
{
"content": {
"parts": [
{
"functionCall": {
"name": "headroom_retrieve",
"args": {"hash": "aaaaaaaaaaaaaaaaaaaaaaaa"},
}
},
{"functionCall": {"name": "client_tool", "args": {}}},
]
}
}
]
}
handler = NativeGeminiHandler([FakeResponse(json_data=response_json)])
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 200
assert len(handler.sent_bodies) == 1
assert json.loads(response.body) == response_json
@pytest.mark.asyncio
async def test_gemini_native_ccr_non_ccr_function_call_is_not_intercepted(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
response_json = {
"candidates": [
{"content": {"parts": [{"functionCall": {"name": "client_tool", "args": {}}}]}}
]
}
handler = NativeGeminiHandler([FakeResponse(json_data=response_json)])
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 200
assert len(handler.sent_bodies) == 1
assert response.body == b"{}"
@pytest.mark.asyncio
async def test_gemini_native_ccr_continuation_error_preserves_upstream_response(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
handler = NativeGeminiHandler(
[
native_ccr_response(),
FakeResponse(status_code=503, content=b"busy", headers={"retry-after": "2"}),
]
)
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 503
assert response.body == b"busy"
assert response.headers["retry-after"] == "2"
@pytest.mark.asyncio
async def test_gemini_native_ccr_continuation_non_json_preserves_upstream_response(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
handler = NativeGeminiHandler(
[native_ccr_response(), FakeResponse(status_code=200, content=b"upstream")]
)
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 200
assert response.body == b"upstream"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"original_content",
[[{"type": "code", "text": "print('x')"}], "plain text", {"key": "value"}, 42],
ids=["code-aware-array", "kompress-text", "mcp-object", "mcp-scalar"],
)
async def test_gemini_native_ccr_uses_real_retrieval_result_shape(
monkeypatch: pytest.MonkeyPatch, original_content
) -> None: # noqa: ANN001
install_native_gemini_compression(monkeypatch)
entry = CompressionEntry(
hash="a" * 24,
original_content=json.dumps(original_content),
compressed_content="compressed",
original_tokens=10,
compressed_tokens=2,
original_item_count=1,
compressed_item_count=1,
tool_name="headroom_retrieve",
tool_call_id="headroom_retrieve",
query_context=None,
created_at=0,
)
class Store:
def get_entry_status(self, hash_key, clean_expired=True): # noqa: ANN001, ARG002
return {"status": "available", "default_ttl_seconds": 1800}
def retrieve(self, hash_key): # noqa: ANN001, ARG002
return entry
monkeypatch.setattr(response_handler_module, "get_compression_store", lambda: Store())
handler = NativeGeminiHandler(
[
native_ccr_response(),
FakeResponse(json_data={"candidates": [{"content": {"parts": [{"text": "done"}]}}]}),
]
)
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 200
function_response = handler.sent_bodies[1]["contents"][-1]["parts"][0]["functionResponse"]
assert function_response["response"]["original_content"] == json.dumps(original_content)
@pytest.mark.asyncio
async def test_gemini_native_ccr_preserves_non_ccr_response(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_native_gemini_compression(monkeypatch)
handler = NativeGeminiHandler([FakeResponse(status_code=503, content=b"busy")])
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 503
assert response.body == b"busy"
@pytest.mark.asyncio
async def test_gemini_native_ccr_residual(monkeypatch: pytest.MonkeyPatch) -> None:
install_native_gemini_compression(monkeypatch)
from headroom.ccr.response_handler import CCRToolResult
handler = NativeGeminiHandler([native_ccr_response()] * 4)
handler.ccr_response_handler._execute_retrieval = lambda call: CCRToolResult(
"headroom_retrieve", "still unresolved", True, 0
)
response = await handler.handle_gemini_generate_content(
FakeRequest(
json.dumps(native_gemini_request()),
headers={"content-type": "application/json"},
path="/v1beta/models/gemini-2.5-flash:generateContent",
),
"gemini-2.5-flash",
)
assert response.status_code == 502
def install_batch_support_modules(
monkeypatch: pytest.MonkeyPatch,
*,
injector_result=None, # noqa: ANN001
tokenizer_count: int = 10,
) -> None:
class FakeInjector:
def __init__(self, **kwargs) -> None: # noqa: ANN003
self.kwargs = kwargs
def process_request(self, messages, tools): # noqa: ANN001, ANN201
if injector_result is not None:
return injector_result
return messages, tools, False
class FakeTokenizer:
def count_messages(self, messages) -> int: # noqa: ANN001
return tokenizer_count
monkeypatch.setitem(sys.modules, "headroom.ccr", SimpleNamespace(CCRToolInjector=FakeInjector))
monkeypatch.setitem(
sys.modules,
"headroom.tokenizers",
SimpleNamespace(get_tokenizer=lambda model: FakeTokenizer()),
)
monkeypatch.setitem(
sys.modules,
"headroom.utils",
SimpleNamespace(extract_user_query=lambda messages: "query"),
)
@pytest.mark.asyncio
async def test_compress_batch_jsonl_without_optimization_handles_invalid_lines(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(monkeypatch, tokenizer_count=12)
handler = DummyBatchHandler()
content = "\n".join(
[
json.dumps(
{"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}}
),
json.dumps({"body": {"model": "gpt-4o", "messages": []}}),
"not-json",
]
)
lines, stats = await handler._compress_batch_jsonl(content, "req-1")
assert len(lines) == 3
assert json.loads(lines[0])["body"]["messages"][0]["content"] == "hi"
assert lines[2] == "not-json"
assert stats == {
"total_requests": 3,
"total_original_tokens": 12,
"total_compressed_tokens": 12,
"total_tokens_saved": 0,
"savings_percent": 0.0,
"errors": 1,
}
@pytest.mark.asyncio
async def test_compress_batch_jsonl_handles_non_object_lines(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# A JSONL line that is valid JSON but not a request object (array/string/
# null), or a request whose `body` isn't a dict, must pass through instead
# of crashing the whole batch (`.get` on a non-dict raises AttributeError,
# which the JSONDecodeError guard does not catch).
install_batch_support_modules(monkeypatch, tokenizer_count=12)
handler = DummyBatchHandler()
content = "\n".join(
[
json.dumps(
{"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}}
),
json.dumps([1, 2, 3]),
json.dumps("hello"),
"null",
json.dumps({"body": "not-a-dict"}),
]
)
lines, stats = await handler._compress_batch_jsonl(content, "req-1")
assert len(lines) == 5
assert json.loads(lines[1]) == [1, 2, 3]
assert json.loads(lines[2]) == "hello"
assert json.loads(lines[3]) is None
assert json.loads(lines[4]) == {"body": "not-a-dict"}
assert stats["total_requests"] == 5
# None of these are JSON decode errors, so the error counter stays at 0.
assert stats["errors"] == 0
@pytest.mark.asyncio
async def test_compress_batch_jsonl_uses_pipeline_and_ccr_injection(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(
monkeypatch,
injector_result=(
[{"role": "system", "content": "compressed"}],
[{"name": "retrieval"}],
True,
),
)
handler = DummyBatchHandler()
handler.config.optimize = True
handler.config.ccr_inject_tool = True
handler.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: SimpleNamespace(
messages=[{"role": "assistant", "content": "short"}],
tokens_before=100,
tokens_after=40,
)
)
lines, stats = await handler._compress_batch_jsonl(
json.dumps(
{
"body": {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"tools": [{"name": "existing"}],
}
}
),
"req-2",
)
body = json.loads(lines[0])["body"]
assert body["messages"] == [{"role": "system", "content": "compressed"}]
assert body["tools"] == [{"name": "retrieval"}]
assert stats["total_tokens_saved"] == 60
assert stats["savings_percent"] == 60.0
@pytest.mark.asyncio
async def test_compress_batch_jsonl_falls_back_when_pipeline_raises(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(monkeypatch, tokenizer_count=33)
handler = DummyBatchHandler()
handler.config.optimize = True
handler.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: (_ for _ in ()).throw(RuntimeError("boom"))
)
lines, stats = await handler._compress_batch_jsonl(
json.dumps({"body": {"messages": [{"role": "user", "content": "hello"}]}}),
"req-3",
)
assert json.loads(lines[0])["body"]["messages"][0]["content"] == "hello"
assert stats["total_original_tokens"] == 33
assert stats["total_compressed_tokens"] == 33
@pytest.mark.asyncio
async def test_batch_passthrough_forwards_request_and_strips_response_headers() -> None:
handler = DummyBatchHandler()
handler.http_client.post_response = FakeResponse(
content=b'{"ok":true}',
headers={"content-encoding": "gzip", "content-length": "20", "x-kept": "1"},
)
response = await handler._batch_passthrough(
FakeRequest(
'{"input_file_id":"file-1"}', headers={"host": "example", "content-length": "10"}
),
{"input_file_id": "file-1"},
)
assert response.status_code == 200
assert dict(response.headers)["x-kept"] == "1"
assert "content-encoding" not in dict(response.headers)
assert handler.http_client.posts[0]["url"] == "https://openai.example/v1/batches"
@pytest.mark.asyncio
async def test_handle_batch_create_validates_json_and_required_fields(
monkeypatch: pytest.MonkeyPatch,
) -> None:
handler = DummyBatchHandler()
async def raise_bad_json(request): # noqa: ANN001
raise ValueError("bad json")
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", raise_bad_json)
bad = await handler.handle_batch_create(FakeRequest("{}"))
assert bad.status_code == 400
assert bad.body.decode().find("invalid_json") > 0
async def missing_file_payload(request): # noqa: ANN001
return {"endpoint": "/v1/chat/completions"}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", missing_file_payload)
missing_file = await handler.handle_batch_create(FakeRequest("{}"))
assert missing_file.status_code == 400
assert missing_file.body.decode().find("input_file_id is required") > 0
async def missing_endpoint_payload(request): # noqa: ANN001
return {"input_file_id": "file-1"}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", missing_endpoint_payload)
missing_endpoint = await handler.handle_batch_create(FakeRequest("{}"))
assert missing_endpoint.status_code == 400
assert missing_endpoint.body.decode().find("endpoint is required") > 0
@pytest.mark.asyncio
async def test_handle_batch_create_passthrough_and_download_failure(
monkeypatch: pytest.MonkeyPatch,
) -> None:
handler = DummyBatchHandler()
passthrough_response = SimpleNamespace(marker="passthrough")
async def fake_passthrough(request, body): # noqa: ANN001
return passthrough_response
monkeypatch.setattr(handler, "_batch_passthrough", fake_passthrough)
async def passthrough_payload(request): # noqa: ANN001
return {"input_file_id": "file-1", "endpoint": "/v1/responses"}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", passthrough_payload)
assert await handler.handle_batch_create(FakeRequest("{}")) is passthrough_response
async def download_missing_payload(request): # noqa: ANN001
return {"input_file_id": "file-1", "endpoint": "/v1/chat/completions"}
async def missing_download(file_id, headers): # noqa: ANN001
return None
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", download_missing_payload)
monkeypatch.setattr(handler, "_download_openai_file", missing_download)
missing = await handler.handle_batch_create(FakeRequest("{}"))
assert missing.status_code == 404
assert missing.body.decode().find("file_not_found") > 0
@pytest.mark.asyncio
async def test_handle_batch_create_handles_empty_upload_failure_and_success(
monkeypatch: pytest.MonkeyPatch,
) -> None:
handler = DummyBatchHandler()
async def request_payload(request): # noqa: ANN001
return {
"input_file_id": "file-1",
"endpoint": "/v1/chat/completions",
"completion_window": "12h",
"metadata": {"source": "test"},
}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", request_payload)
async def fake_download(file_id, headers): # noqa: ANN001
return "downloaded"
monkeypatch.setattr(handler, "_download_openai_file", fake_download)
async def empty_compress(content, request_id): # noqa: ANN001
return [], {
"total_requests": 0,
"total_original_tokens": 0,
"total_compressed_tokens": 0,
"total_tokens_saved": 0,
"savings_percent": 0.0,
"errors": 0,
}
monkeypatch.setattr(handler, "_compress_batch_jsonl", empty_compress)
empty = await handler.handle_batch_create(FakeRequest("{}"))
assert empty.status_code == 400
assert empty.body.decode().find("empty_file") > 0
async def compressed(content, request_id): # noqa: ANN001
return ['{"body":{}}'], {
"total_requests": 1,
"total_original_tokens": 20,
"total_compressed_tokens": 10,
"total_tokens_saved": 10,
"savings_percent": 50.0,
"errors": 0,
}
monkeypatch.setattr(handler, "_compress_batch_jsonl", compressed)
async def upload_failed_file(content, filename, headers): # noqa: ANN001
return None
monkeypatch.setattr(handler, "_upload_openai_file", upload_failed_file)
upload_failed = await handler.handle_batch_create(FakeRequest("{}"))
assert upload_failed.status_code == 500
assert upload_failed.body.decode().find("upload_failed") > 0
handler.http_client.post_response = FakeResponse(
content=b'{"id":"batch_123","object":"batch"}',
headers={"content-encoding": "gzip", "content-length": "12", "x-openai": "1"},
)
async def upload_success(content, filename, headers): # noqa: ANN001
return "file-compressed"
monkeypatch.setattr(handler, "_upload_openai_file", upload_success)
success = await handler.handle_batch_create(
FakeRequest(
"{}", headers={"host": "proxy", "content-length": "4", "authorization": "Bearer test"}
)
)
assert success.status_code == 200
success_headers = dict(success.headers)
assert success_headers["x-headroom-tokens-saved"] == "10"
assert success_headers["x-headroom-savings-percent"] == "50.0"
assert success_headers["x-openai"] == "1"
# PR-A3: byte-faithful forwarder writes ``content`` (raw bytes), not
# ``json``. Round-trip the captured bytes back to a dict for assertion.
last_post = handler.http_client.posts[-1]
if "json" in last_post:
sent_body = last_post["json"]
else:
sent_body = json.loads(last_post["content"].decode("utf-8"))
assert sent_body["metadata"]["headroom_compressed"] == "true"
assert sent_body["metadata"]["headroom_original_file_id"] == "file-1"
assert handler.metrics.record_calls[-1]["provider"] == "openai"
@pytest.mark.asyncio
async def test_handle_batch_create_records_failure_on_exception(
monkeypatch: pytest.MonkeyPatch,
) -> None:
handler = DummyBatchHandler()
async def request_payload(request): # noqa: ANN001
return {"input_file_id": "file-1", "endpoint": "/v1/chat/completions"}
async def boom(file_id, headers): # noqa: ANN001
raise RuntimeError("boom")
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", request_payload)
monkeypatch.setattr(handler, "_download_openai_file", boom)
response = await handler.handle_batch_create(FakeRequest("{}"))
assert response.status_code == 500
assert handler.metrics.failed_calls == [{"provider": "batch"}]
@pytest.mark.asyncio
async def test_download_and_upload_openai_file_helpers() -> None:
handler = DummyBatchHandler()
handler.http_client.get_response = FakeResponse(status_code=200, text="jsonl-content")
downloaded = await handler._download_openai_file("file-1", {"authorization": "Bearer token"})
assert downloaded == "jsonl-content"
assert handler.http_client.gets[0]["url"] == "https://openai.example/v1/files/file-1/content"
handler.http_client.get_response = FakeResponse(status_code=404, text="missing")
assert await handler._download_openai_file("file-2", {}) is None
handler.http_client.post_response = FakeResponse(
status_code=200,
json_data={"id": "file-uploaded"},
headers={"content-type": "application/json"},
)
file_id = await handler._upload_openai_file(
'{"body":{}}',
"compressed.jsonl",
{"authorization": "Bearer token", "content-type": "application/json"},
)
assert file_id == "file-uploaded"
post_call = handler.http_client.posts[-1]
assert post_call["headers"] == {"authorization": "Bearer token"}
assert post_call["files"]["file"][0] == "compressed.jsonl"
handler.http_client.post_response = FakeResponse(status_code=500, text="fail")
assert await handler._upload_openai_file("{}", "bad.jsonl", {}) is None
handler.http_client.raise_post = RuntimeError("network")
assert await handler._upload_openai_file("{}", "bad.jsonl", {}) is None
@pytest.mark.asyncio
async def test_store_google_batch_context_persists_transformed_requests(
monkeypatch: pytest.MonkeyPatch,
) -> None:
stored_contexts: list[object] = []
class FakeBatchContext:
def __init__(self, **kwargs) -> None: # noqa: ANN003
self.kwargs = kwargs
self.requests: list[object] = []
def add_request(self, request) -> None: # noqa: ANN001
self.requests.append(request)
class FakeBatchRequestContext:
def __init__(self, **kwargs) -> None: # noqa: ANN003
self.kwargs = kwargs
class FakeStore:
async def store(self, context) -> None: # noqa: ANN001
stored_contexts.append(context)
monkeypatch.setitem(
sys.modules,
"headroom.ccr",
SimpleNamespace(
BatchContext=FakeBatchContext,
BatchRequestContext=FakeBatchRequestContext,
get_batch_context_store=lambda: FakeStore(),
),
)
handler = DummyBatchHandler()
await handler._store_google_batch_context(
"batches/123",
[
{
"metadata": {"key": "req-1"},
"request": {
"contents": [{"parts": [{"text": "hello"}]}],
"systemInstruction": {"parts": [{"text": "system"}]},
"tools": [{"name": "tool"}],
},
}
],
"gemini-2.0",
"api-key",
)
context = stored_contexts[0]
assert context.kwargs["batch_id"] == "batches/123"
assert context.requests[0].kwargs["custom_id"] == "req-1"
assert context.requests[0].kwargs["messages"] == [{"role": "user", "content": "hello"}]
assert context.requests[0].kwargs["system_instruction"] == "system"
@pytest.mark.asyncio
async def test_handle_google_batch_results_passes_through_early_exit_cases(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class FakeStore:
async def get(self, batch_name): # noqa: ANN001
return None
monkeypatch.setitem(
sys.modules,
"headroom.ccr",
SimpleNamespace(
BatchResultProcessor=lambda http_client: None,
get_batch_context_store=lambda: FakeStore(),
),
)
handler = DummyBatchHandler()
request = FakeRequest(
"{}", headers={"x-goog-api-key": "secret"}, method="GET", path="/v1beta/batches/b1"
)
handler.http_client.get_response = FakeResponse(
status_code=500, content=b"bad", headers={"x-upstream": "1"}
)
error_response = await handler.handle_google_batch_results(request, "batches/b1")
assert error_response.status_code == 500
assert dict(error_response.headers)["x-upstream"] == "1"
class BadJsonResponse(FakeResponse):
def json(self): # noqa: ANN201
raise json.JSONDecodeError("bad", "x", 0)
handler.http_client.get_response = BadJsonResponse(
status_code=200, content=b"plain", headers={"x-upstream": "2"}
)
non_json = await handler.handle_google_batch_results(request, "batches/b1")
assert non_json.status_code == 200
assert dict(non_json.headers)["x-upstream"] == "2"
handler.http_client.get_response = FakeResponse(
status_code=200,
content=b"{}",
json_data={"metadata": {"state": "RUNNING"}},
)
running = await handler.handle_google_batch_results(request, "batches/b1")
assert running.status_code == 200
handler.http_client.get_response = FakeResponse(
status_code=200,
content=b"{}",
json_data={"metadata": {"state": "SUCCEEDED"}, "response": {"responses": []}},
)
no_results = await handler.handle_google_batch_results(request, "batches/b1")
assert no_results.status_code == 200
handler.http_client.get_response = FakeResponse(
status_code=200,
content=b"{}",
json_data={"metadata": {"state": "SUCCEEDED"}, "response": {"responses": [{"id": 1}]}},
)
handler.config.ccr_inject_tool = False
no_ccr = await handler.handle_google_batch_results(request, "batches/b1")
assert no_ccr.status_code == 200
assert "key=secret" in handler.http_client.gets[-1]["url"]
@pytest.mark.asyncio
async def test_handle_google_batch_results_processes_completed_results(
monkeypatch: pytest.MonkeyPatch,
) -> None:
processed_calls: list[tuple[str, list[object], str]] = []
class FakeProcessed:
def __init__(
self, result, custom_id: str, was_processed: bool, continuation_rounds: int
) -> None: # noqa: ANN001
self.result = result
self.custom_id = custom_id
self.was_processed = was_processed
self.continuation_rounds = continuation_rounds
class FakeProcessor:
def __init__(self, http_client) -> None: # noqa: ANN001
self.http_client = http_client
async def process_results(self, batch_name, results, provider): # noqa: ANN001
processed_calls.append((batch_name, results, provider))
return [
FakeProcessed({"id": "processed"}, "req-1", True, 2),
FakeProcessed({"id": "unchanged"}, "req-2", False, 0),
]
class FakeStore:
async def get(self, batch_name): # noqa: ANN001
return SimpleNamespace(batch_name=batch_name)
monkeypatch.setitem(
sys.modules,
"headroom.ccr",
SimpleNamespace(
BatchResultProcessor=FakeProcessor,
get_batch_context_store=lambda: FakeStore(),
),
)
handler = DummyBatchHandler()
handler.config.ccr_inject_tool = True
handler.http_client.get_response = FakeResponse(
status_code=200,
content=b"{}",
json_data={
"metadata": {"state": "SUCCEEDED"},
"response": {"responses": [{"id": "raw-1"}, {"id": "raw-2"}]},
},
)
response = await handler.handle_google_batch_results(
FakeRequest("{}", method="GET", path="/v1beta/batches/b1"),
"batches/b1",
)
payload = json.loads(response.body)
assert payload["response"]["responses"] == [{"id": "processed"}, {"id": "unchanged"}]
assert processed_calls == [("batches/b1", [{"id": "raw-1"}, {"id": "raw-2"}], "google")]
assert handler.metrics.record_calls[-1]["model"] == "batch:ccr-processed"
@pytest.mark.asyncio
async def test_google_batch_passthrough_helpers_forward_and_track_metrics() -> None:
handler = DummyBatchHandler()
handler.http_client.post_response = FakeResponse(
content=b'{"ok":true}',
headers={"content-encoding": "gzip", "content-length": "10", "x-kept": "1"},
)
handler.http_client.post_response = FakeResponse(
content=b'{"ok":true}',
headers={"content-encoding": "gzip", "content-length": "10", "x-kept": "1"},
)
passthrough = await handler._google_batch_passthrough(
FakeRequest(
"body", headers={"host": "proxy", "content-length": "4", "x-goog-api-key": "secret"}
),
"gemini-pro",
{"batch": {}},
)
assert passthrough.status_code == 200
assert dict(passthrough.headers)["x-kept"] == "1"
assert "key=secret" in handler.http_client.posts[-1]["url"]
assert handler.metrics.record_calls[-1]["model"] == "passthrough:batch:gemini-pro"
handler.http_client.get_response = FakeResponse(
content=b'{"state":"ok"}',
headers={"content-encoding": "gzip", "content-length": "10", "x-kept": "2"},
)
response = await handler.handle_google_batch_passthrough(
FakeRequest(
"ping",
headers={"host": "proxy", "x-goog-api-key": "secret"},
method="DELETE",
path="/v1beta/batches/b1",
query="alt=json",
),
"b1",
)
assert response.status_code == 200
assert dict(response.headers)["x-kept"] == "2"
get_call = handler.http_client.requests[-1]
assert get_call["url"] == "https://gemini.example/v1beta/batches/b1?alt=json&key=secret"
assert handler.metrics.record_calls[-1]["model"] == "passthrough:batches"
@pytest.mark.asyncio
async def test_handle_google_batch_create_validates_and_passthroughs(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(monkeypatch)
handler = DummyBatchHandler()
too_large = await handler.handle_google_batch_create(
FakeRequest("{}", headers={"content-length": str(200 * 1024 * 1024)}),
"gemini-pro",
)
assert too_large.status_code == 413
async def bad_json(request): # noqa: ANN001
raise ValueError("bad json")
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", bad_json)
invalid = await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
assert invalid.status_code == 400
passthrough_response = SimpleNamespace(kind="passthrough")
async def fake_google_passthrough(request, model, body=None): # noqa: ANN001
return passthrough_response
async def no_inline(request): # noqa: ANN001
return {"batch": {"input_config": {"requests": {"requests": []}}}}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", no_inline)
monkeypatch.setattr(handler, "_google_batch_passthrough", fake_google_passthrough)
assert (
await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
is passthrough_response
)
@pytest.mark.asyncio
async def test_handle_google_batch_create_success_and_failure_paths(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(monkeypatch)
handler = DummyBatchHandler()
handler.config.optimize = True
handler.config.ccr_inject_tool = True
handler.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: SimpleNamespace(
messages=[{"role": "user", "content": "compressed"}],
timing={"compress": 1.2},
tokens_before=100,
tokens_after=40,
)
)
class FakeInjector:
def __init__(self, **kwargs) -> None: # noqa: ANN003
pass
def process_request(self, messages, tools): # noqa: ANN001, ANN201
return (
messages + [{"role": "system", "content": "retrieval"}],
[{"name": "retrieval"}],
True,
)
monkeypatch.setitem(sys.modules, "headroom.ccr", SimpleNamespace(CCRToolInjector=FakeInjector))
stored: list[tuple[str, list[dict[str, object]], str, str | None]] = []
async def fake_store(batch_name, requests_list, model, api_key): # noqa: ANN001
stored.append((batch_name, requests_list, model, api_key))
async def fake_retry(method, url, headers, body, **kwargs): # noqa: ANN001
return FakeResponse(
status_code=200,
content=b'{"name":"batches/123"}',
headers={"content-encoding": "gzip", "content-length": "10", "x-upstream": "1"},
json_data={"name": "batches/123"},
)
async def good_payload(request): # noqa: ANN001
return {
"batch": {
"input_config": {
"requests": {
"requests": [
{
"request": {
"contents": [{"parts": [{"text": "hello"}]}],
"tools": [{"functionDeclarations": [{"name": "existing"}]}],
},
"metadata": {"key": "req-1"},
}
]
}
}
}
}
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", good_payload)
monkeypatch.setattr(handler, "_retry_request", fake_retry)
monkeypatch.setattr(handler, "_store_google_batch_context", fake_store)
response = await handler.handle_google_batch_create(
FakeRequest("{}", headers={"x-goog-api-key": "secret"}),
"gemini-pro",
)
assert response.status_code == 200
assert dict(response.headers)["x-upstream"] == "1"
assert handler.metrics.record_calls[-1]["provider"] == "google"
assert handler.metrics.record_calls[-1]["tokens_saved"] == 60
assert stored[0][0] == "batches/123"
assert stored[0][2:] == ("gemini-pro", "secret")
assert stored[0][1][0]["metadata"] == {"key": "req-1"}
async def broken_retry(method, url, headers, body, **kwargs): # noqa: ANN001
raise RuntimeError("forward failed")
monkeypatch.setattr(handler, "_retry_request", broken_retry)
failed = await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
assert failed.status_code == 500
@pytest.mark.asyncio
async def test_handle_google_batch_create_covers_passthrough_revert_and_store_failures(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(
monkeypatch, injector_result=([{"role": "user", "content": "kept"}], None, False)
)
handler = DummyBatchHandler()
handler.config.optimize = True
handler.config.ccr_inject_tool = True
pipeline_calls: list[dict[str, object]] = []
handler.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: (
pipeline_calls.append(kwargs)
or SimpleNamespace(
messages=[{"role": "user", "content": "inflated"}],
timing={},
tokens_before=40,
tokens_after=80,
)
)
)
def fake_to_messages(contents, system_instruction): # noqa: ANN001, ANN201
if contents and "inlineData" in contents[0]["parts"][0]:
return ([{"role": "user", "content": "binary"}], [0])
return ([{"role": "user", "content": "compress"}], [])
def fake_to_gemini(messages): # noqa: ANN001, ANN201
return ([{"parts": [{"text": "new"}]}], {"parts": [{"text": "sys"}]})
async def payload(request): # noqa: ANN001
return {
"batch": {
"input_config": {
"requests": {
"requests": [
{"request": {"contents": []}, "metadata": {"key": "empty"}},
{
"request": {"contents": [{"parts": [{"inlineData": "x"}]}]},
"metadata": {"key": "preserved"},
},
{
"request": {
"contents": [{"parts": [{"text": "hello"}]}],
"tools": [
{"other": True},
{"functionDeclarations": [{"name": "existing"}]},
],
},
"metadata": {"key": "optimized"},
},
]
}
}
}
}
seen_bodies: list[dict[str, object]] = []
async def retry(method, url, headers, body, **kwargs): # noqa: ANN001
seen_bodies.append(body)
return FakeResponse(status_code=200, content=b"{}", json_data={"name": "batches/123"})
async def broken_store(batch_name, requests_list, model, api_key): # noqa: ANN001
raise RuntimeError("store failed")
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", payload)
monkeypatch.setattr(handler, "_gemini_contents_to_messages", fake_to_messages)
monkeypatch.setattr(handler, "_messages_to_gemini_contents", fake_to_gemini)
monkeypatch.setattr(handler, "_retry_request", retry)
monkeypatch.setattr(handler, "_store_google_batch_context", broken_store)
response = await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
assert response.status_code == 200
assert len(pipeline_calls) == 1
assert handler.metrics.record_calls[-1]["tokens_saved"] == 0
assert (
seen_bodies[0]["batch"]["input_config"]["requests"]["requests"][0]["metadata"]["key"]
== "empty"
)
optimized = seen_bodies[0]["batch"]["input_config"]["requests"]["requests"][2]["request"]
assert optimized["contents"][0] == {"parts": [{"text": "new"}]}
assert optimized["systemInstruction"] == {"parts": [{"text": "sys"}]}
@pytest.mark.asyncio
async def test_handle_google_batch_create_preserves_functioncall_response_order(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A batch request that interleaves text turns with text-less
functionCall/functionResponse entries must reach Google with all entries
intact and in order. The old raw-index restore loop overwrote the model's
answer with the functionCall and dropped the functionResponse."""
class RealConvHandler(batch_module.BatchHandlerMixin, GeminiHandlerMixin):
# Real Gemini converters + _rebuild_gemini_contents (no stubs), so the
# actual index interleaving runs.
GEMINI_API_URL = "https://gemini.example"
def __init__(self) -> None:
self.http_client = FakeHttpClient()
self.metrics = FakeMetrics()
self.config = SimpleNamespace(
optimize=True, ccr_inject_tool=False, ccr_inject_system_instructions=False
)
self.openai_provider = SimpleNamespace(get_context_limit=lambda m: 8192)
# No-op pipeline: return the messages unchanged, no token inflation.
self.openai_pipeline = SimpleNamespace(
apply=lambda **kw: SimpleNamespace(
messages=kw["messages"], timing={}, tokens_before=100, tokens_after=100
)
)
self.captured_body: dict | None = None
async def _next_request_id(self) -> str:
return "req-1"
async def _record_request_outcome(self, outcome) -> None: # noqa: ANN001
pass
def _extract_tags(self, headers: dict) -> dict[str, str]:
return {}
async def _run_compression_in_executor(self, fn, *, timeout): # noqa: ANN001, ANN201
return fn()
async def _store_google_batch_context(self, *a, **k) -> None: # noqa: ANN002, ANN003
pass
async def _retry_request(self, method, url, headers, body, **kwargs): # noqa: ANN001, ANN201
# Capture the (in-place mutated) forwarded batch body for assertions.
self.captured_body = body
return FakeResponse(status_code=200, content=b"{}", json_data={"name": "batches/1"})
handler = RealConvHandler()
contents = [
{"role": "user", "parts": [{"text": "What's the weather in Paris?"}]},
{
"role": "model",
"parts": [{"functionCall": {"name": "get_weather", "args": {"city": "Paris"}}}],
},
{
"role": "user",
"parts": [{"functionResponse": {"name": "get_weather", "response": {"temp_c": 18}}}],
},
{"role": "model", "parts": [{"text": "It's 18C and cloudy in Paris."}]},
]
batch_body = {
"batch": {
"input_config": {
"requests": {"requests": [{"request": {"contents": contents}, "metadata": {}}]}
}
}
}
async def payload(request): # noqa: ANN001, ANN201
return batch_body
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", payload)
resp = await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
assert resp.status_code == 200
out = handler.captured_body["batch"]["input_config"]["requests"]["requests"][0]["request"][
"contents"
]
# All four entries survive in order. The old loop produced only two, dropping
# the functionResponse and overwriting the model answer with the functionCall.
assert len(out) == 4
assert "text" in out[0]["parts"][0]
assert out[1]["parts"][0].get("functionCall", {}).get("name") == "get_weather"
assert out[2]["parts"][0].get("functionResponse", {}).get("name") == "get_weather"
assert "Paris" in out[3]["parts"][0]["text"]
@pytest.mark.asyncio
async def test_handle_google_batch_create_preserves_sibling_tools(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A batch request whose tools array carries googleSearch / codeExecution
alongside functionDeclarations must reach Google with those siblings intact.
The old code collapsed the whole array to a single functionDeclarations
entry, silently disabling Google Search and code execution."""
class RealConvHandler(batch_module.BatchHandlerMixin, GeminiHandlerMixin):
GEMINI_API_URL = "https://gemini.example"
def __init__(self) -> None:
self.http_client = FakeHttpClient()
self.metrics = FakeMetrics()
self.config = SimpleNamespace(
optimize=True, ccr_inject_tool=False, ccr_inject_system_instructions=False
)
self.openai_provider = SimpleNamespace(get_context_limit=lambda m: 8192)
self.openai_pipeline = SimpleNamespace(
apply=lambda **kw: SimpleNamespace(
messages=kw["messages"], timing={}, tokens_before=100, tokens_after=100
)
)
self.captured_body: dict | None = None
async def _next_request_id(self) -> str:
return "req-1"
async def _record_request_outcome(self, outcome) -> None: # noqa: ANN001
pass
def _extract_tags(self, headers: dict) -> dict[str, str]:
return {}
async def _run_compression_in_executor(self, fn, *, timeout): # noqa: ANN001, ANN201
return fn()
async def _store_google_batch_context(self, *a, **k) -> None: # noqa: ANN002, ANN003
pass
async def _retry_request(self, method, url, headers, body, **kwargs): # noqa: ANN001, ANN201
self.captured_body = body
return FakeResponse(status_code=200, content=b"{}", json_data={"name": "batches/1"})
handler = RealConvHandler()
tools = [
{"functionDeclarations": [{"name": "get_weather"}]},
{"googleSearch": {}},
{"codeExecution": {}},
]
batch_body = {
"batch": {
"input_config": {
"requests": {
"requests": [
{
"request": {
"contents": [{"role": "user", "parts": [{"text": "hello there"}]}],
"tools": tools,
},
"metadata": {},
}
]
}
}
}
}
async def payload(request): # noqa: ANN001, ANN201
return batch_body
monkeypatch.setattr("headroom.proxy.helpers._read_request_json", payload)
resp = await handler.handle_google_batch_create(FakeRequest("{}"), "gemini-pro")
assert resp.status_code == 200
out_tools = handler.captured_body["batch"]["input_config"]["requests"]["requests"][0][
"request"
]["tools"]
keys = [next(iter(entry)) for entry in out_tools]
assert "googleSearch" in keys
assert "codeExecution" in keys
assert "functionDeclarations" in keys
@pytest.mark.asyncio
async def test_google_batch_passthrough_without_body_and_query_variants() -> None:
handler = DummyBatchHandler()
handler.http_client.post_response = FakeResponse(content=b"ok", headers={"x-upstream": "1"})
response = await handler._google_batch_passthrough(
FakeRequest("raw-body", headers={"host": "proxy"}, method="POST"),
"gemini-pro",
)
assert response.status_code == 200
assert handler.http_client.posts[-1]["content"] == b"raw-body"
handler.http_client.get_response = FakeResponse(content=b"{}", headers={"x-upstream": "2"})
passthrough = await handler.handle_google_batch_passthrough(
FakeRequest(
"{}",
headers={"host": "proxy", "x-goog-api-key": "secret"},
method="GET",
path="/v1beta/batches/b1",
),
"b1",
)
assert passthrough.status_code == 200
assert (
handler.http_client.requests[-1]["url"]
== "https://gemini.example/v1beta/batches/b1?key=secret"
)
@pytest.mark.asyncio
async def test_batch_helper_methods_and_openai_file_error_branches() -> None:
handler = DummyBatchHandler()
marker = object()
async def fake_passthrough(request, base_url): # noqa: ANN001
return marker
handler.handle_passthrough = fake_passthrough
request = FakeRequest("{}")
assert await handler.handle_batch_list(request) is marker
assert await handler.handle_batch_get(request, "b1") is marker
assert await handler.handle_batch_cancel(request, "b1") is marker
handler.http_client.raise_get = RuntimeError("download boom")
assert await handler._download_openai_file("file-1", {}) is None
handler.http_client.raise_get = None
handler.http_client.post_response = FakeResponse(status_code=200, json_data={})
assert await handler._upload_openai_file("{}", "missing-id.jsonl", {}) is None
@pytest.mark.asyncio
async def test_store_google_batch_context_without_system_text(
monkeypatch: pytest.MonkeyPatch,
) -> None:
stored_contexts: list[object] = []
class FakeBatchContext:
def __init__(self, **kwargs) -> None: # noqa: ANN003
self.kwargs = kwargs
self.requests: list[object] = []
def add_request(self, request) -> None: # noqa: ANN001
self.requests.append(request)
class FakeBatchRequestContext:
def __init__(self, **kwargs) -> None: # noqa: ANN003
self.kwargs = kwargs
class FakeStore:
async def store(self, context) -> None: # noqa: ANN001
stored_contexts.append(context)
handler = DummyBatchHandler()
monkeypatch.setitem(
sys.modules,
"headroom.ccr",
SimpleNamespace(
BatchContext=FakeBatchContext,
BatchRequestContext=FakeBatchRequestContext,
get_batch_context_store=lambda: FakeStore(),
),
)
await handler._store_google_batch_context(
"batches/456",
[
{
"request": {
"contents": [{"parts": [{"text": "hello"}]}],
"systemInstruction": {"parts": ["bad"]},
}
}
],
"gemini-2.0",
None,
)
context = stored_contexts[0]
assert context.kwargs["api_key"] is None
assert context.requests[0].kwargs["custom_id"] == ""
assert context.requests[0].kwargs["system_instruction"] is None
@pytest.mark.asyncio
async def test_compress_batch_jsonl_skips_blank_lines_and_preserves_tools_when_not_injected(
monkeypatch: pytest.MonkeyPatch,
) -> None:
install_batch_support_modules(
monkeypatch,
injector_result=([{"role": "assistant", "content": "short"}], [{"name": "orig"}], False),
)
handler = DummyBatchHandler()
handler.config.optimize = True
handler.config.ccr_inject_tool = True
handler.openai_pipeline = SimpleNamespace(
apply=lambda **kwargs: SimpleNamespace(
messages=[{"role": "assistant", "content": "short"}],
tokens_before=50,
tokens_after=10,
)
)
lines, stats = await handler._compress_batch_jsonl(
"\n"
+ json.dumps(
{
"body": {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
"tools": [{"name": "orig"}],
}
}
)
+ "\n",
"req-extra",
)
assert len(lines) == 1
body = json.loads(lines[0])["body"]
assert body["tools"] == [{"name": "orig"}]
assert stats["total_requests"] == 1
assert stats["errors"] == 0