From 4bdc1482af32b9cfc6a6d0dd5d4af3ad7110b28e Mon Sep 17 00:00:00 2001 From: Kayzo Date: Fri, 17 Apr 2026 17:45:08 +0000 Subject: [PATCH 1/2] feat: add Pi/Codex and Cloud Code Assist compatibility routes - Adds /v1/codex/responses aliases for OpenAI Codex clients configured with /v1 base URLs - Uses JWT-derived account routing for /v1/responses/* subpaths (compact, cancel, etc) - Adds /v1internal:streamGenerateContent aliases for Cloud Code Assist / Antigravity - Preserves upstream HTTP error status/body in StreamingResponse (fixes empty SSE drops) --- docs/content/docs/proxy.mdx | 18 ++ headroom/proxy/handlers/gemini.py | 155 +++++++++++++++++- headroom/proxy/handlers/streaming.py | 15 +- headroom/proxy/server.py | 30 +++- tests/test_proxy_codex_route_aliases.py | 78 ++++++++- ...st_proxy_google_cloudcode_route_aliases.py | 91 ++++++++++ .../test_proxy_streaming_ratelimit_headers.py | 39 +++++ 7 files changed, 420 insertions(+), 6 deletions(-) create mode 100644 tests/test_proxy_google_cloudcode_route_aliases.py diff --git a/docs/content/docs/proxy.mdx b/docs/content/docs/proxy.mdx index d1313f784..cf72774fc 100644 --- a/docs/content/docs/proxy.mdx +++ b/docs/content/docs/proxy.mdx @@ -135,6 +135,24 @@ Anthropic API format. The proxy compresses messages, forwards to Anthropic, and OpenAI API format. The proxy compresses messages, forwards to OpenAI, and returns the response. +### `POST /v1/responses` + +OpenAI Responses API format. The proxy compresses `input` payloads where applicable, forwards the request, and returns the response. + +For Codex-compatible clients, the proxy also accepts these alias paths and routes them through the same handler: +- `POST /v1/codex/responses` +- `POST /backend-api/responses` +- `POST /backend-api/codex/responses` + +Matching WebSocket and subpath aliases are also supported for Codex flows. + +### `POST /v1internal:streamGenerateContent` + +Google Cloud Code Assist / Antigravity compatibility endpoint used by Pi-style `google-gemini-cli` and `google-antigravity` providers. + +The proxy also accepts: +- `POST /v1/v1internal:streamGenerateContent` + ### `POST /v1/compress` Compression-only endpoint. Compresses messages without calling any LLM. Used by the TypeScript SDK. diff --git a/headroom/proxy/handlers/gemini.py b/headroom/proxy/handlers/gemini.py index 1f26bc46f..724cf7ad5 100644 --- a/headroom/proxy/handlers/gemini.py +++ b/headroom/proxy/handlers/gemini.py @@ -9,7 +9,7 @@ import json import logging import os import time -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from fastapi import Request @@ -17,10 +17,35 @@ if TYPE_CHECKING: logger = logging.getLogger("headroom.proxy") +DEFAULT_GEMINI_API_URL = "https://generativelanguage.googleapis.com" +DEFAULT_CLOUDCODE_API_URL = "https://cloudcode-pa.googleapis.com" +ANTIGRAVITY_DAILY_API_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com" +ANTIGRAVITY_AUTOPUSH_API_URL = "https://autopush-cloudcode-pa.sandbox.googleapis.com" + class GeminiHandlerMixin: """Mixin providing Gemini API handler methods for HeadroomProxy.""" + def _is_cloudcode_antigravity_request( + self, body: dict[str, Any], headers: dict[str, str] + ) -> bool: + """Detect Pi/OpenClaw antigravity requests routed via Cloud Code Assist.""" + user_agent = headers.get("user-agent", "") + return ( + body.get("requestType") == "agent" + or body.get("userAgent") == "antigravity" + or user_agent.startswith("antigravity/") + ) + + def _resolve_cloudcode_base_url(self, is_antigravity: bool) -> str: + """Resolve upstream base URL for Pi Cloud Code Assist / Antigravity traffic.""" + configured = self.GEMINI_API_URL.rstrip("/") + if configured and configured != DEFAULT_GEMINI_API_URL: + return configured + if is_antigravity: + return ANTIGRAVITY_DAILY_API_URL + return DEFAULT_CLOUDCODE_API_URL + def _has_non_text_parts(self, content: dict) -> bool: """Check if a Gemini content entry has non-text parts. @@ -447,6 +472,134 @@ class GeminiHandlerMixin: }, ) + async def handle_google_cloudcode_stream( + self, + request: Request, + ) -> StreamingResponse | JSONResponse: + """Handle Pi/OpenClaw Google Cloud Code Assist and Antigravity streaming requests.""" + from fastapi.responses import JSONResponse + + from headroom.proxy.helpers import _read_request_json + from headroom.tokenizers import get_tokenizer + from headroom.utils import extract_user_query + + start_time = time.time() + request_id = await self._next_request_id() + + try: + body = await _read_request_json(request) + except (json.JSONDecodeError, ValueError) as e: + return JSONResponse( + status_code=400, + content={ + "error": { + "message": f"Invalid request body: {e!s}", + "code": 400, + } + }, + ) + + request_payload = body.get("request") + if not isinstance(request_payload, dict): + return JSONResponse( + status_code=400, + content={ + "error": { + "message": "Invalid Cloud Code Assist request: missing request payload", + "code": 400, + } + }, + ) + + model = body.get("model", "unknown") + contents = request_payload.get("contents", []) + headers = dict(request.headers.items()) + headers.pop("host", None) + headers.pop("content-length", None) + headers.pop("accept-encoding", None) + tags = self._extract_tags(headers) + is_antigravity = self._is_cloudcode_antigravity_request(body, headers) + + system_instruction = request_payload.get("systemInstruction") + optimization_system_instruction = None if is_antigravity else system_instruction + messages, preserved_indices = self._gemini_contents_to_messages( + contents if isinstance(contents, list) else [], optimization_system_instruction + ) + preserved_contents = { + idx: contents[idx] + for idx in preserved_indices + if isinstance(contents, list) and idx < len(contents) + } + + tokenizer = get_tokenizer(model) + original_tokens = tokenizer.count_messages(messages) if messages else 0 + optimized_messages = messages + optimized_tokens = original_tokens + transforms_applied: list[str] = [] + + _license_ok = self.usage_reporter.should_compress if self.usage_reporter else True + if self.config.optimize and messages and _license_ok: + try: + context_limit = self.openai_provider.get_context_limit(model) + result = self.openai_pipeline.apply( + messages=messages, + model=model, + model_limit=context_limit, + context=extract_user_query(messages), + ) + if result.messages != messages: + optimized_messages = result.messages + transforms_applied = result.transforms_applied + original_tokens = result.tokens_before + optimized_tokens = result.tokens_after + except Exception as e: + logger.warning(f"[{request_id}] Cloud Code Assist optimization failed: {e}") + + if optimized_tokens > original_tokens: + logger.warning( + f"[{request_id}] Cloud Code Assist optimization inflated tokens " + f"({original_tokens} -> {optimized_tokens}), reverting to original messages" + ) + optimized_messages = messages + optimized_tokens = original_tokens + transforms_applied = [] + + if optimized_messages != messages: + optimized_contents, optimized_system = self._messages_to_gemini_contents( + optimized_messages + ) + for orig_idx, original_content in preserved_contents.items(): + if orig_idx < len(optimized_contents): + optimized_contents[orig_idx] = original_content + request_payload["contents"] = optimized_contents + if not is_antigravity: + if optimized_system: + request_payload["systemInstruction"] = optimized_system + elif "systemInstruction" in request_payload: + del request_payload["systemInstruction"] + + tokens_saved = original_tokens - optimized_tokens + optimization_latency = (time.time() - start_time) * 1000 + base_url = self._resolve_cloudcode_base_url(is_antigravity) + stream_url = f"{base_url}/v1internal:streamGenerateContent" + if request.url.query: + stream_url = f"{stream_url}?{request.url.query}" + + return await self._stream_response( + stream_url, + headers, + body, + "gemini", + model, + request_id, + original_tokens, + optimized_tokens, + tokens_saved, + transforms_applied, + tags, + optimization_latency, + ) + async def handle_gemini_stream_generate_content( self, request: Request, diff --git a/headroom/proxy/handlers/streaming.py b/headroom/proxy/handlers/streaming.py index 3bb3f0b45..aa397c067 100644 --- a/headroom/proxy/handlers/streaming.py +++ b/headroom/proxy/handlers/streaming.py @@ -462,7 +462,7 @@ class StreamingMixin: 3. Makes continuation requests until no memory tools remain 4. Streams the final response to the client """ - from fastapi.responses import StreamingResponse + from fastapi.responses import Response, StreamingResponse from headroom.proxy.cost import _summarize_transforms from headroom.proxy.helpers import MAX_SSE_BUFFER_SIZE @@ -511,6 +511,19 @@ class StreamingMixin: return StreamingResponse(_error_gen(), media_type="text/event-stream") + if upstream_response.status_code >= 400: + response_headers = dict(upstream_response.headers) + response_headers.pop("content-length", None) + response_headers.pop("transfer-encoding", None) + response_headers.pop("connection", None) + error_content = await upstream_response.aread() + await upstream_response.aclose() + return Response( + content=error_content, + status_code=upstream_response.status_code, + headers=response_headers, + ) + # Forward upstream ratelimit headers to the client forwarded_headers = { k: v for k, v in upstream_response.headers.items() if "ratelimit" in k.lower() diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index 7ee1f19a4..d21f8095b 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -2033,6 +2033,11 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: """OpenAI Responses API (new API introduced March 2025).""" return await proxy.handle_openai_responses(request) + @app.post("/v1/codex/responses") + async def openai_v1_codex_responses(request: Request): + """Pi/OpenAI Codex compatibility path for OpenAI-style /v1 base URLs.""" + return await proxy.handle_openai_responses(request) + @app.post("/backend-api/responses") async def openai_codex_responses(request: Request): """OpenAI Codex Responses API path preserved from ChatGPT backend.""" @@ -2048,6 +2053,11 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: """OpenAI Responses API via WebSocket (Codex gpt-5.4+).""" await proxy.handle_openai_responses_ws(websocket) + @app.websocket("/v1/codex/responses") + async def openai_v1_codex_responses_ws(websocket: WebSocket): + """Pi/OpenAI Codex compatibility WebSocket path for /v1 base URLs.""" + await proxy.handle_openai_responses_ws(websocket) + # OpenAI Responses API sub-endpoints (passthrough). # Codex sub-agents use /v1/responses/compact and other sub-paths # that we don't need to compress — just forward with correct auth routing. @@ -2056,13 +2066,16 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: """Passthrough for /v1/responses/* sub-endpoints (compact, cancel, etc.).""" from fastapi.responses import Response + from headroom.proxy.handlers.openai import _resolve_codex_routing_headers + headers = dict(request.headers.items()) headers.pop("host", None) + headers, is_chatgpt_auth = _resolve_codex_routing_headers(headers) # Route to correct endpoint based on auth mode. # ChatGPT session auth (codex login) uses chatgpt.com with /responses/... # path (no /v1/ prefix). API key auth uses api.openai.com/v1/responses/... - if headers.get("chatgpt-account-id"): + if is_chatgpt_auth: url = f"https://chatgpt.com/backend-api/codex/responses/{sub_path}" else: url = f"{proxy.OPENAI_API_URL}/v1/responses/{sub_path}" @@ -2089,6 +2102,11 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: logger.error(f"Passthrough /v1/responses/{sub_path} failed: {e}") return Response(content=str(e), status_code=502) + @app.api_route("/v1/codex/responses/{sub_path:path}", methods=["GET", "POST", "DELETE"]) + async def openai_v1_codex_responses_sub(request: Request, sub_path: str): + """Passthrough for Pi/OpenAI Codex /v1/codex/responses/* sub-endpoints.""" + return await openai_responses_sub(request, sub_path) + @app.websocket("/backend-api/responses") async def openai_codex_responses_ws(websocket: WebSocket): """OpenAI Codex Responses WebSocket path preserved from ChatGPT backend.""" @@ -2148,6 +2166,16 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: """Gemini countTokens API with compression applied.""" return await proxy.handle_gemini_count_tokens(request, model) + @app.post("/v1internal:streamGenerateContent") + async def google_cloudcode_stream_generate_content(request: Request): + """Google Cloud Code Assist / Antigravity compatibility streaming endpoint.""" + return await proxy.handle_google_cloudcode_stream(request) + + @app.post("/v1/v1internal:streamGenerateContent") + async def google_cloudcode_stream_generate_content_v1(request: Request): + """Compatibility endpoint for clients configured with a /v1 proxy base URL.""" + return await proxy.handle_google_cloudcode_stream(request) + # ========================================================================= # Databricks Native Endpoints # ========================================================================= diff --git a/tests/test_proxy_codex_route_aliases.py b/tests/test_proxy_codex_route_aliases.py index 797361f61..d9dd20d0b 100644 --- a/tests/test_proxy_codex_route_aliases.py +++ b/tests/test_proxy_codex_route_aliases.py @@ -1,3 +1,6 @@ +import base64 +import json + import httpx from fastapi import WebSocket from fastapi.responses import JSONResponse @@ -6,6 +9,16 @@ from fastapi.testclient import TestClient from headroom.proxy.server import HeadroomProxy, ProxyConfig, create_app +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_codex_responses_aliases_delegate_to_openai_handler(monkeypatch): async def fake_handle(self, request): # type: ignore[no-untyped-def] return JSONResponse({"ok": True, "path": request.url.path}) @@ -13,7 +26,11 @@ def test_codex_responses_aliases_delegate_to_openai_handler(monkeypatch): monkeypatch.setattr(HeadroomProxy, "handle_openai_responses", fake_handle) with TestClient(create_app(ProxyConfig())) as client: - for path in ("/backend-api/responses", "/backend-api/codex/responses"): + for path in ( + "/v1/codex/responses", + "/backend-api/responses", + "/backend-api/codex/responses", + ): response = client.post(path, json={"model": "gpt-5.3-codex"}) assert response.status_code == 200 assert response.json() == {"ok": True, "path": path} @@ -31,11 +48,19 @@ def test_codex_responses_websocket_aliases_delegate_to_openai_handler(monkeypatc monkeypatch.setattr(HeadroomProxy, "handle_openai_responses_ws", fake_handle_ws) with TestClient(create_app(ProxyConfig())) as client: - for path in ("/backend-api/responses", "/backend-api/codex/responses"): + for path in ( + "/v1/codex/responses", + "/backend-api/responses", + "/backend-api/codex/responses", + ): with client.websocket_connect(path) as websocket: assert websocket.receive_json() == {"ok": True, "path": path} - assert seen_paths == ["/backend-api/responses", "/backend-api/codex/responses"] + assert seen_paths == [ + "/v1/codex/responses", + "/backend-api/responses", + "/backend-api/codex/responses", + ] def test_codex_responses_subpath_aliases_delegate_to_passthrough(): @@ -55,6 +80,10 @@ def test_codex_responses_subpath_aliases_delegate_to_passthrough(): client.app.state.proxy.http_client = fake_http_client client.app.state.proxy.OPENAI_API_URL = "https://api.openai.test" + pi_response = client.post( + "/v1/codex/responses/compact?trace=0", + json={"model": "gpt-5.3-codex"}, + ) api_key_response = client.post( "/backend-api/responses/compact?trace=1", json={"model": "gpt-5.3-codex"}, @@ -65,9 +94,52 @@ def test_codex_responses_subpath_aliases_delegate_to_passthrough(): json={"model": "gpt-5.3-codex"}, ) + assert pi_response.status_code == 200 assert api_key_response.status_code == 200 assert chatgpt_response.status_code == 200 assert fake_http_client.calls == [ + ("POST", "https://api.openai.test/v1/responses/compact?trace=0"), ("POST", "https://api.openai.test/v1/responses/compact?trace=1"), ("POST", "https://chatgpt.com/backend-api/codex/responses/compact?trace=2"), ] + + +def test_codex_responses_subpath_passthrough_derives_chatgpt_routing_from_jwt(): + class FakeAsyncClient: + def __init__(self) -> None: + self.calls: list[tuple[str, str, dict[str, str]]] = [] + + async def request(self, method, url, **kwargs): # type: ignore[no-untyped-def] + self.calls.append((method, url, dict(kwargs.get("headers", {})))) + return httpx.Response(200, json={"method": method, "url": url}) + + async def aclose(self) -> None: + return None + + token = _jwt( + { + "https://api.openai.com/auth": { + "chatgpt_account_id": "acct-from-jwt", + } + } + ) + + with TestClient(create_app(ProxyConfig())) as client: + fake_http_client = FakeAsyncClient() + client.app.state.proxy.http_client = fake_http_client + client.app.state.proxy.OPENAI_API_URL = "https://api.openai.test" + + response = client.post( + "/v1/codex/responses/compact?trace=jwt", + headers={"Authorization": f"Bearer {token}"}, + json={"model": "gpt-5.4"}, + ) + + assert response.status_code == 200 + assert len(fake_http_client.calls) == 1 + + method, url, headers = fake_http_client.calls[0] + assert method == "POST" + assert url == "https://chatgpt.com/backend-api/codex/responses/compact?trace=jwt" + assert headers["authorization"] == f"Bearer {token}" + assert headers["ChatGPT-Account-ID"] == "acct-from-jwt" diff --git a/tests/test_proxy_google_cloudcode_route_aliases.py b/tests/test_proxy_google_cloudcode_route_aliases.py new file mode 100644 index 000000000..a90595fa6 --- /dev/null +++ b/tests/test_proxy_google_cloudcode_route_aliases.py @@ -0,0 +1,91 @@ +from fastapi.responses import JSONResponse +from fastapi.testclient import TestClient + +from headroom.proxy.server import HeadroomProxy, ProxyConfig, create_app + +CLOUDCODE_BODY = { + "project": "test-project", + "model": "gemini-3.1-pro-high", + "userAgent": "pi-coding-agent", + "request": { + "contents": [ + { + "role": "user", + "parts": [{"text": "Reply with pong."}], + } + ] + }, +} + +ANTIGRAVITY_BODY = { + "project": "test-project", + "model": "claude-sonnet-4-6", + "requestType": "agent", + "userAgent": "antigravity", + "request": { + "contents": [ + { + "role": "user", + "parts": [{"text": "Reply with pong."}], + } + ] + }, +} + + +def test_google_cloudcode_alias_routes_delegate_to_handler(monkeypatch): + async def fake_handle(self, request): # type: ignore[no-untyped-def] + return JSONResponse({"ok": True, "path": request.url.path}) + + monkeypatch.setattr(HeadroomProxy, "handle_google_cloudcode_stream", fake_handle) + + with TestClient(create_app(ProxyConfig())) as client: + for path in ( + "/v1internal:streamGenerateContent", + "/v1/v1internal:streamGenerateContent", + ): + response = client.post(path, params={"alt": "sse"}, json=CLOUDCODE_BODY) + assert response.status_code == 200 + assert response.json() == {"ok": True, "path": path} + + +def test_antigravity_cloudcode_route_uses_daily_endpoint(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + with TestClient(create_app(ProxyConfig(optimize=False))) as client: + response = client.post( + "/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=ANTIGRAVITY_BODY, + ) + + assert response.status_code == 200 + assert response.json() == { + "url": "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:streamGenerateContent?alt=sse", + "provider": "gemini", + "model": "claude-sonnet-4-6", + } + + +def test_cloudcode_route_uses_default_cloudcode_endpoint(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + with TestClient(create_app(ProxyConfig(optimize=False))) as client: + response = client.post( + "/v1/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=CLOUDCODE_BODY, + ) + + assert response.status_code == 200 + assert response.json() == { + "url": "https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent?alt=sse", + "provider": "gemini", + "model": "gemini-3.1-pro-high", + } diff --git a/tests/test_proxy_streaming_ratelimit_headers.py b/tests/test_proxy_streaming_ratelimit_headers.py index 6c1d306b3..8af5e8ff0 100644 --- a/tests/test_proxy_streaming_ratelimit_headers.py +++ b/tests/test_proxy_streaming_ratelimit_headers.py @@ -68,6 +68,7 @@ class TestStreamingRatelimitHeaderForwarding: if extra_headers: headers.update(extra_headers) mock_response.headers = httpx.Headers(headers) + mock_response.status_code = 200 # Simulate a simple SSE stream sse_data = ( @@ -198,6 +199,44 @@ class TestStreamingRatelimitHeaderForwarding: # No ratelimit headers to forward assert result.headers.get("anthropic-ratelimit-tokens-limit") is None + @pytest.mark.asyncio + async def test_upstream_http_error_preserves_status_and_body(self): + """Upstream non-200 streaming responses should preserve HTTP status/body.""" + proxy = self._create_mock_proxy() + mock_response = self._create_mock_upstream_response() + mock_response.status_code = 503 + mock_response.aread = AsyncMock( + return_value=b'{"error":{"message":"capacity exhausted"}}' + ) + mock_response.aclose = AsyncMock() + + mock_request = MagicMock() + proxy.http_client.build_request = MagicMock(return_value=mock_request) + proxy.http_client.send = AsyncMock(return_value=mock_response) + + result = await proxy._stream_response( + url="https://api.anthropic.com/v1/messages", + headers={"x-api-key": "sk-test"}, + body={ + "model": "claude-sonnet-4-20250514", + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + provider="anthropic", + model="claude-sonnet-4-20250514", + request_id="test-http-error", + original_tokens=10, + optimized_tokens=10, + tokens_saved=0, + transforms_applied=[], + tags={}, + optimization_latency=0.0, + ) + + assert result.status_code == 503 + assert result.body == b'{"error":{"message":"capacity exhausted"}}' + @pytest.mark.asyncio async def test_connect_error_returns_sse_error(self): """Connection errors should return an SSE error event (not crash).""" From 9f0a813004e0f915939f0d89b469a91c68407381 Mon Sep 17 00:00:00 2001 From: Kayzo Date: Fri, 17 Apr 2026 23:09:12 +0000 Subject: [PATCH 2/2] fix: address PR review feedback --- headroom/cli/proxy.py | 16 +- headroom/proxy/handlers/gemini.py | 12 +- headroom/proxy/handlers/openai.py | 2 + headroom/proxy/handlers/streaming.py | 269 +++++++++++------- headroom/proxy/models.py | 1 + headroom/proxy/server.py | 27 +- tests/test_proxy_codex_route_aliases.py | 22 +- ...st_proxy_google_cloudcode_route_aliases.py | 107 +++++++ .../test_proxy_streaming_ratelimit_headers.py | 74 ++++- 9 files changed, 412 insertions(+), 118 deletions(-) diff --git a/headroom/cli/proxy.py b/headroom/cli/proxy.py index a3863595e..28932fb9a 100644 --- a/headroom/cli/proxy.py +++ b/headroom/cli/proxy.py @@ -159,6 +159,11 @@ from .main import main default=None, help="Custom Gemini API URL for passthrough endpoints (env: GEMINI_TARGET_API_URL)", ) +@click.option( + "--cloudcode-api-url", + default=None, + help="Custom Cloud Code Assist API URL for compatibility endpoints (env: CLOUDCODE_TARGET_API_URL)", +) @click.option( "--region", default="us-west-2", @@ -216,6 +221,7 @@ def proxy( anthropic_api_url: str | None, openai_api_url: str | None, gemini_api_url: str | None, + cloudcode_api_url: str | None, region: str, bedrock_region: str | None, bedrock_profile: str | None, @@ -250,6 +256,7 @@ def proxy( effective_anthropic_api_url = anthropic_api_url or os.environ.get("ANTHROPIC_TARGET_API_URL") effective_openai_api_url = openai_api_url or os.environ.get("OPENAI_TARGET_API_URL") effective_gemini_api_url = gemini_api_url or os.environ.get("GEMINI_TARGET_API_URL") + effective_cloudcode_api_url = cloudcode_api_url or os.environ.get("CLOUDCODE_TARGET_API_URL") # Resolve anyllm provider: env var takes precedence over CLI default (matches argparse path) effective_anyllm_provider = os.environ.get("HEADROOM_ANYLLM_PROVIDER") or anyllm_provider @@ -284,6 +291,7 @@ def proxy( anthropic_api_url=effective_anthropic_api_url, openai_api_url=effective_openai_api_url, gemini_api_url=effective_gemini_api_url, + cloudcode_api_url=effective_cloudcode_api_url, mode=effective_mode, optimize=not no_optimize, cache_enabled=not no_cache, @@ -335,6 +343,7 @@ def proxy( anthropic_url = config.anthropic_api_url or "https://api.anthropic.com" openai_url = config.openai_api_url or "https://api.openai.com" + cloudcode_url = config.cloudcode_api_url or "https://cloudcode-pa.googleapis.com" backend_section = "" if config.backend == "anyllm" or config.backend.startswith("anyllm-"): @@ -415,9 +424,10 @@ Starting proxy server... {stateless_line}{telemetry_line} {backend_section} Routing: - /v1/messages → {anthropic_url} - /v1/chat/completions → {openai_url} - /v1/responses → {openai_url} (HTTP + WebSocket) + /v1/messages → {anthropic_url} + /v1/chat/completions → {openai_url} + /v1/responses → {openai_url} (HTTP + WebSocket) + /v1internal:streamGenerateContent → {cloudcode_url} Usage: Claude Code: ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude diff --git a/headroom/proxy/handlers/gemini.py b/headroom/proxy/handlers/gemini.py index 724cf7ad5..4aa246ba9 100644 --- a/headroom/proxy/handlers/gemini.py +++ b/headroom/proxy/handlers/gemini.py @@ -17,10 +17,8 @@ if TYPE_CHECKING: logger = logging.getLogger("headroom.proxy") -DEFAULT_GEMINI_API_URL = "https://generativelanguage.googleapis.com" DEFAULT_CLOUDCODE_API_URL = "https://cloudcode-pa.googleapis.com" ANTIGRAVITY_DAILY_API_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com" -ANTIGRAVITY_AUTOPUSH_API_URL = "https://autopush-cloudcode-pa.sandbox.googleapis.com" class GeminiHandlerMixin: @@ -30,21 +28,19 @@ class GeminiHandlerMixin: self, body: dict[str, Any], headers: dict[str, str] ) -> bool: """Detect Pi/OpenClaw antigravity requests routed via Cloud Code Assist.""" - user_agent = headers.get("user-agent", "") + user_agent = headers.get("user-agent", "").lower() + body_user_agent = str(body.get("userAgent", "")).lower() return ( body.get("requestType") == "agent" - or body.get("userAgent") == "antigravity" + or body_user_agent == "antigravity" or user_agent.startswith("antigravity/") ) def _resolve_cloudcode_base_url(self, is_antigravity: bool) -> str: """Resolve upstream base URL for Pi Cloud Code Assist / Antigravity traffic.""" - configured = self.GEMINI_API_URL.rstrip("/") - if configured and configured != DEFAULT_GEMINI_API_URL: - return configured if is_antigravity: return ANTIGRAVITY_DAILY_API_URL - return DEFAULT_CLOUDCODE_API_URL + return getattr(self, "CLOUDCODE_API_URL", DEFAULT_CLOUDCODE_API_URL).rstrip("/") def _has_non_text_parts(self, content: dict) -> bool: """Check if a Gemini content entry has non-text parts. diff --git a/headroom/proxy/handlers/openai.py b/headroom/proxy/handlers/openai.py index b631c7b26..134829884 100644 --- a/headroom/proxy/handlers/openai.py +++ b/headroom/proxy/handlers/openai.py @@ -42,6 +42,8 @@ def _decode_openai_bearer_payload(headers: dict[str, str]) -> dict[str, Any] | N payload = token.split(".", 2)[1] payload += "=" * (-len(payload) % 4) + # Intentionally no signature verification here: this is only a best-effort + # routing hint extractor. Upstream still performs the actual auth/authz checks. try: decoded = base64.urlsafe_b64decode(payload.encode("ascii")) data = json.loads(decoded.decode("utf-8")) diff --git a/headroom/proxy/handlers/streaming.py b/headroom/proxy/handlers/streaming.py index aa397c067..4700b33af 100644 --- a/headroom/proxy/handlers/streaming.py +++ b/headroom/proxy/handlers/streaming.py @@ -13,7 +13,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from fastapi.responses import StreamingResponse + from fastapi.responses import Response, StreamingResponse import httpx @@ -432,6 +432,116 @@ class StreamingMixin: except Exception as e: logger.debug(f"[{request_id}] CCR Feedback recording failed: {e}") + async def _finalize_stream_response( + self, + *, + body: dict, + provider: str, + model: str, + request_id: str, + original_tokens: int, + optimized_tokens: int, + tokens_saved: int, + transforms_applied: list[str], + optimization_latency: float, + stream_state: dict[str, Any], + start_time: float, + pipeline_timing: dict[str, float] | None = None, + prefix_tracker: Any | None = None, + original_messages: list[dict] | None = None, + full_sse_data: str = "", + parsed_response: dict[str, Any] | None = None, + ) -> None: + from headroom.proxy.cost import _summarize_transforms + + total_latency = (time.time() - start_time) * 1000 + output_tokens = stream_state["output_tokens"] + if output_tokens is None: + output_tokens = stream_state["total_bytes"] // 40 + logger.warning( + f"[{request_id}] Could not parse output_tokens from SSE, " + f"estimating {output_tokens} from {stream_state['total_bytes']} bytes" + ) + + cache_read_tokens = stream_state["cache_read_input_tokens"] or 0 + cache_write_tokens = stream_state["cache_creation_input_tokens"] or 0 + cache_write_5m_tokens = stream_state["cache_creation_ephemeral_5m_input_tokens"] or 0 + cache_write_1h_tokens = stream_state["cache_creation_ephemeral_1h_input_tokens"] or 0 + uncached_input_tokens = max(optimized_tokens - cache_read_tokens - cache_write_tokens, 0) + + num_msgs = len(body.get("messages", [])) + cache_hit_pct = ( + round(cache_read_tokens / (cache_read_tokens + cache_write_tokens) * 100) + if (cache_read_tokens + cache_write_tokens) > 0 + else 0 + ) + logger.info( + f"[{request_id}] PERF " + f"model={model} msgs={num_msgs} " + f"tok_before={original_tokens} tok_after={optimized_tokens} " + f"tok_saved={tokens_saved} " + f"cache_read={cache_read_tokens} cache_write={cache_write_tokens} " + f"cache_hit_pct={cache_hit_pct} " + f"opt_ms={optimization_latency:.0f} " + f"transforms={_summarize_transforms(transforms_applied)}" + ) + + if prefix_tracker is not None: + import copy as _copy + + forwarded_messages = body.get("messages", []) + next_forwarded = _copy.deepcopy(forwarded_messages) + next_original = _copy.deepcopy(original_messages or forwarded_messages) + + if full_sse_data and provider == "anthropic": + _parsed = ( + parsed_response + if parsed_response is not None + else self._parse_sse_to_response(full_sse_data, provider) + ) + if _parsed: + asst_msg = self._assistant_message_from_response_json(_parsed) + if asst_msg is not None: + next_forwarded.append(_copy.deepcopy(asst_msg)) + next_original.append(_copy.deepcopy(asst_msg)) + + prefix_tracker.update_from_response( + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + messages=next_forwarded, + original_messages=next_original, + ) + + if self.cost_tracker: + self.cost_tracker.record_tokens( + model, + tokens_saved, + optimized_tokens, + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + cache_write_5m_tokens=cache_write_5m_tokens, + cache_write_1h_tokens=cache_write_1h_tokens, + uncached_tokens=uncached_input_tokens, + ) + + if getattr(self, "metrics", None) is not None: + await self.metrics.record_request( + provider=provider, + model=model, + input_tokens=optimized_tokens, + output_tokens=output_tokens, + tokens_saved=tokens_saved, + latency_ms=total_latency, + overhead_ms=optimization_latency, + ttfb_ms=stream_state["ttfb_ms"] or total_latency, + pipeline_timing=pipeline_timing, + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + cache_write_5m_tokens=cache_write_5m_tokens, + cache_write_1h_tokens=cache_write_1h_tokens, + uncached_input_tokens=uncached_input_tokens, + ) + async def _stream_response( self, url: str, @@ -450,7 +560,7 @@ class StreamingMixin: pipeline_timing: dict[str, float] | None = None, prefix_tracker: Any | None = None, original_messages: list[dict] | None = None, - ) -> StreamingResponse: + ) -> Response | StreamingResponse: """Stream response with metrics tracking and memory tool handling. Parses SSE events to extract actual usage information from the API response @@ -464,7 +574,6 @@ class StreamingMixin: """ from fastapi.responses import Response, StreamingResponse - from headroom.proxy.cost import _summarize_transforms from headroom.proxy.helpers import MAX_SSE_BUFFER_SIZE start_time = time.time() @@ -512,12 +621,57 @@ class StreamingMixin: return StreamingResponse(_error_gen(), media_type="text/event-stream") if upstream_response.status_code >= 400: + logger.warning( + "[%s] Forwarding upstream streaming error status=%s url=%s", + request_id, + upstream_response.status_code, + url, + ) response_headers = dict(upstream_response.headers) response_headers.pop("content-length", None) response_headers.pop("transfer-encoding", None) response_headers.pop("connection", None) - error_content = await upstream_response.aread() - await upstream_response.aclose() + response_headers.pop("content-encoding", None) + + try: + error_content = await upstream_response.aread() + except Exception as read_error: + logger.warning( + "[%s] Failed reading upstream error body status=%s url=%s error=%s", + request_id, + upstream_response.status_code, + url, + read_error, + ) + error_content = json.dumps( + { + "error": { + "message": "Failed to read upstream error response body", + "details": str(read_error), + } + } + ).encode("utf-8") + response_headers["content-type"] = "application/json" + finally: + await upstream_response.aclose() + + stream_state["total_bytes"] = len(error_content) + await self._finalize_stream_response( + body=body, + provider=provider, + model=model, + request_id=request_id, + original_tokens=original_tokens, + optimized_tokens=optimized_tokens, + tokens_saved=tokens_saved, + transforms_applied=transforms_applied, + optimization_latency=optimization_latency, + stream_state=stream_state, + start_time=start_time, + pipeline_timing=pipeline_timing, + prefix_tracker=prefix_tracker, + original_messages=original_messages, + ) return Response( content=error_content, status_code=upstream_response.status_code, @@ -675,102 +829,23 @@ class StreamingMixin: } yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode() finally: - # Record metrics after stream completes - total_latency = (time.time() - start_time) * 1000 - - # Use actual output tokens from API if available, otherwise estimate - output_tokens = stream_state["output_tokens"] - if output_tokens is None: - # Fallback: estimate from bytes (but this is inaccurate for SSE) - # Use a more conservative estimate - SSE overhead is ~10-20x - output_tokens = stream_state["total_bytes"] // 40 - logger.debug( - f"[{request_id}] No usage in stream, estimated {output_tokens} output tokens" - ) - - # Use optimized_tokens for dashboard metrics (what we actually sent). - # API's input_tokens is the non-cached portion only, which is - # misleading for aggregation (often just 1 with prompt caching). - cache_read_tokens = stream_state["cache_read_input_tokens"] - cache_write_tokens = stream_state["cache_creation_input_tokens"] - cache_write_5m_tokens = stream_state["cache_creation_ephemeral_5m_input_tokens"] - cache_write_1h_tokens = stream_state["cache_creation_ephemeral_1h_input_tokens"] - uncached_input_tokens = stream_state.get("input_tokens") or 0 - - # Structured perf log line for `headroom perf` analysis - num_msgs = len(body.get("messages", [])) - cache_hit_pct = ( - round(cache_read_tokens / (cache_read_tokens + cache_write_tokens) * 100) - if (cache_read_tokens + cache_write_tokens) > 0 - else 0 - ) - logger.info( - f"[{request_id}] PERF " - f"model={model} msgs={num_msgs} " - f"tok_before={original_tokens} tok_after={optimized_tokens} " - f"tok_saved={tokens_saved} " - f"cache_read={cache_read_tokens} cache_write={cache_write_tokens} " - f"cache_hit_pct={cache_hit_pct} " - f"opt_ms={optimization_latency:.0f} " - f"transforms={_summarize_transforms(transforms_applied)}" - ) - - # Update prefix cache tracker for next turn (streaming path) - if prefix_tracker is not None: - import copy as _copy - - forwarded_messages = body.get("messages", []) - next_forwarded = _copy.deepcopy(forwarded_messages) - next_original = _copy.deepcopy(original_messages or forwarded_messages) - - # Reconstruct assistant response from SSE data so the - # prefix tracker accounts for it in the cached prefix - if full_sse_data and provider == "anthropic": - _parsed = ( - parsed_response - if parsed_response is not None - else self._parse_sse_to_response(full_sse_data, provider) - ) - if _parsed: - asst_msg = self._assistant_message_from_response_json(_parsed) - if asst_msg is not None: - next_forwarded.append(_copy.deepcopy(asst_msg)) - next_original.append(_copy.deepcopy(asst_msg)) - - prefix_tracker.update_from_response( - cache_read_tokens=cache_read_tokens, - cache_write_tokens=cache_write_tokens, - messages=next_forwarded, - original_messages=next_original, - ) - - if self.cost_tracker: - self.cost_tracker.record_tokens( - model, - tokens_saved, - optimized_tokens, - cache_read_tokens=cache_read_tokens, - cache_write_tokens=cache_write_tokens, - cache_write_5m_tokens=cache_write_5m_tokens, - cache_write_1h_tokens=cache_write_1h_tokens, - uncached_tokens=uncached_input_tokens, - ) - - await self.metrics.record_request( + await self._finalize_stream_response( + body=body, provider=provider, model=model, - input_tokens=optimized_tokens, # What we sent, not API's non-cached count - output_tokens=output_tokens, + request_id=request_id, + original_tokens=original_tokens, + optimized_tokens=optimized_tokens, tokens_saved=tokens_saved, - latency_ms=total_latency, - overhead_ms=optimization_latency, - ttfb_ms=stream_state["ttfb_ms"] or 0, + transforms_applied=transforms_applied, + optimization_latency=optimization_latency, + stream_state=stream_state, + start_time=start_time, pipeline_timing=pipeline_timing, - cache_read_tokens=cache_read_tokens, - cache_write_tokens=cache_write_tokens, - cache_write_5m_tokens=cache_write_5m_tokens, - cache_write_1h_tokens=cache_write_1h_tokens, - uncached_input_tokens=uncached_input_tokens, + prefix_tracker=prefix_tracker, + original_messages=original_messages, + full_sse_data=full_sse_data, + parsed_response=parsed_response, ) return StreamingResponse( diff --git a/headroom/proxy/models.py b/headroom/proxy/models.py index 446bae552..04a74aa6f 100644 --- a/headroom/proxy/models.py +++ b/headroom/proxy/models.py @@ -79,6 +79,7 @@ class ProxyConfig: anthropic_api_url: str | None = None # Custom Anthropic API URL override openai_api_url: str | None = None # Custom OpenAI API URL override gemini_api_url: str | None = None # Custom Gemini API URL override + cloudcode_api_url: str | None = None # Custom Cloud Code Assist API URL override # Backend: "anthropic" (direct API), "litellm-*" (via LiteLLM), or "anyllm" (via any-llm) backend: str = "anthropic" diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index d21f8095b..6e4e2cadb 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -203,32 +203,49 @@ class HeadroomProxy( ANTHROPIC_API_URL = "https://api.anthropic.com" OPENAI_API_URL = "https://api.openai.com" GEMINI_API_URL = "https://generativelanguage.googleapis.com" + CLOUDCODE_API_URL = "https://cloudcode-pa.googleapis.com" def __init__(self, config: ProxyConfig): self.config = config self.config.mode = normalize_proxy_mode(self.config.mode) - # Override ANTHROPIC_API_URL with config if set - # Strip trailing /v1 or /v1/ to avoid double-path (e.g., .../v1/v1/models) + # Reset per-instance API targets first so test runs and multiple app instances + # do not leak class-level overrides across each other. + HeadroomProxy.ANTHROPIC_API_URL = "https://api.anthropic.com" + HeadroomProxy.OPENAI_API_URL = "https://api.openai.com" + HeadroomProxy.GEMINI_API_URL = "https://generativelanguage.googleapis.com" + HeadroomProxy.CLOUDCODE_API_URL = "https://cloudcode-pa.googleapis.com" + + # Override ANTHROPIC_API_URL with config if set. + # Strip trailing /v1 or /v1/ to avoid double-path (e.g., .../v1/v1/models). if config.anthropic_api_url: url = config.anthropic_api_url.rstrip("/") if url.endswith("/v1"): url = url[:-3] HeadroomProxy.ANTHROPIC_API_URL = url - # Override OPENAI_API_URL with config if set - # Strip trailing /v1 or /v1/ to avoid double-path (e.g., .../v1/v1/models) + # Override OPENAI_API_URL with config if set. + # Strip trailing /v1 or /v1/ to avoid double-path (e.g., .../v1/v1/models). if config.openai_api_url: url = config.openai_api_url.rstrip("/") if url.endswith("/v1"): url = url[:-3] HeadroomProxy.OPENAI_API_URL = url - # Override GEMINI_API_URL with config if set + # Override GEMINI_API_URL with config if set. if config.gemini_api_url: gurl = config.gemini_api_url.rstrip("/") + if gurl.endswith("/v1"): + gurl = gurl[:-3] HeadroomProxy.GEMINI_API_URL = gurl + # Override CLOUDCODE_API_URL with config if set. + if config.cloudcode_api_url: + curl = config.cloudcode_api_url.rstrip("/") + if curl.endswith("/v1"): + curl = curl[:-3] + HeadroomProxy.CLOUDCODE_API_URL = curl + # Initialize providers self.anthropic_provider = AnthropicProvider() self.openai_provider = OpenAIProvider() diff --git a/tests/test_proxy_codex_route_aliases.py b/tests/test_proxy_codex_route_aliases.py index d9dd20d0b..d74529d2e 100644 --- a/tests/test_proxy_codex_route_aliases.py +++ b/tests/test_proxy_codex_route_aliases.py @@ -2,6 +2,7 @@ import base64 import json import httpx +import pytest from fastapi import WebSocket from fastapi.responses import JSONResponse from fastapi.testclient import TestClient @@ -104,7 +105,22 @@ def test_codex_responses_subpath_aliases_delegate_to_passthrough(): ] -def test_codex_responses_subpath_passthrough_derives_chatgpt_routing_from_jwt(): +@pytest.mark.parametrize( + ("path", "expected_url"), + [ + ( + "/v1/codex/responses/compact?trace=jwt", + "https://chatgpt.com/backend-api/codex/responses/compact?trace=jwt", + ), + ( + "/v1/responses/compact?trace=jwt-old", + "https://chatgpt.com/backend-api/codex/responses/compact?trace=jwt-old", + ), + ], +) +def test_codex_responses_subpath_passthrough_derives_chatgpt_routing_from_jwt( + path, expected_url +): class FakeAsyncClient: def __init__(self) -> None: self.calls: list[tuple[str, str, dict[str, str]]] = [] @@ -130,7 +146,7 @@ def test_codex_responses_subpath_passthrough_derives_chatgpt_routing_from_jwt(): client.app.state.proxy.OPENAI_API_URL = "https://api.openai.test" response = client.post( - "/v1/codex/responses/compact?trace=jwt", + path, headers={"Authorization": f"Bearer {token}"}, json={"model": "gpt-5.4"}, ) @@ -140,6 +156,6 @@ def test_codex_responses_subpath_passthrough_derives_chatgpt_routing_from_jwt(): method, url, headers = fake_http_client.calls[0] assert method == "POST" - assert url == "https://chatgpt.com/backend-api/codex/responses/compact?trace=jwt" + assert url == expected_url assert headers["authorization"] == f"Bearer {token}" assert headers["ChatGPT-Account-ID"] == "acct-from-jwt" diff --git a/tests/test_proxy_google_cloudcode_route_aliases.py b/tests/test_proxy_google_cloudcode_route_aliases.py index a90595fa6..5f2a85592 100644 --- a/tests/test_proxy_google_cloudcode_route_aliases.py +++ b/tests/test_proxy_google_cloudcode_route_aliases.py @@ -89,3 +89,110 @@ def test_cloudcode_route_uses_default_cloudcode_endpoint(monkeypatch): "provider": "gemini", "model": "gemini-3.1-pro-high", } + + +def test_cloudcode_route_uses_cloudcode_api_override(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + with TestClient( + create_app( + ProxyConfig(optimize=False, cloudcode_api_url="https://cloudcode-proxy.test/v1") + ) + ) as client: + response = client.post( + "/v1/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=CLOUDCODE_BODY, + ) + + assert response.status_code == 200 + assert response.json() == { + "url": "https://cloudcode-proxy.test/v1internal:streamGenerateContent?alt=sse", + "provider": "gemini", + "model": "gemini-3.1-pro-high", + } + + +def test_antigravity_header_detection_is_case_insensitive(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + body = { + **CLOUDCODE_BODY, + "model": "claude-opus-4-6-thinking", + } + + with TestClient(create_app(ProxyConfig(optimize=False))) as client: + response = client.post( + "/v1internal:streamGenerateContent", + params={"alt": "sse"}, + headers={"User-Agent": "Antigravity/1.2.3 Darwin/arm64"}, + json=body, + ) + + assert response.status_code == 200 + assert response.json() == { + "url": "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:streamGenerateContent?alt=sse", + "provider": "gemini", + "model": "claude-opus-4-6-thinking", + } + + +def test_antigravity_route_does_not_cross_route_to_cloudcode_override(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + with TestClient( + create_app( + ProxyConfig(optimize=False, cloudcode_api_url="https://cloudcode-proxy.test") + ) + ) as client: + response = client.post( + "/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=ANTIGRAVITY_BODY, + ) + + assert response.status_code == 200 + assert response.json() == { + "url": "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:streamGenerateContent?alt=sse", + "provider": "gemini", + "model": "claude-sonnet-4-6", + } + + +def test_cloudcode_override_does_not_leak_between_app_instances(monkeypatch): + async def fake_stream(self, url, _headers, _body, provider, model, *_args, **_kwargs): # type: ignore[no-untyped-def] + return JSONResponse({"url": url, "provider": provider, "model": model}) + + monkeypatch.setattr(HeadroomProxy, "_stream_response", fake_stream) + + with TestClient( + create_app( + ProxyConfig(optimize=False, cloudcode_api_url="https://cloudcode-proxy.test") + ) + ) as client: + first = client.post( + "/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=CLOUDCODE_BODY, + ) + + with TestClient(create_app(ProxyConfig(optimize=False))) as client: + second = client.post( + "/v1internal:streamGenerateContent", + params={"alt": "sse"}, + json=CLOUDCODE_BODY, + ) + + assert first.status_code == 200 + assert first.json()["url"] == "https://cloudcode-proxy.test/v1internal:streamGenerateContent?alt=sse" + assert second.status_code == 200 + assert second.json()["url"] == "https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent?alt=sse" diff --git a/tests/test_proxy_streaming_ratelimit_headers.py b/tests/test_proxy_streaming_ratelimit_headers.py index 8af5e8ff0..392139195 100644 --- a/tests/test_proxy_streaming_ratelimit_headers.py +++ b/tests/test_proxy_streaming_ratelimit_headers.py @@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest +import headroom.proxy.handlers.streaming as streaming_module from headroom.proxy.server import HeadroomProxy @@ -27,6 +28,8 @@ class TestStreamingRatelimitHeaderForwarding: proxy.cost_tracker = MagicMock() proxy.cost_tracker.estimate_cost.return_value = 0.001 proxy.cost_tracker.record_request.return_value = None + proxy.metrics = MagicMock() + proxy.metrics.record_request = AsyncMock() proxy.stats = { "requests_total": 0, "requests_optimized": 0, @@ -40,6 +43,7 @@ class TestStreamingRatelimitHeaderForwarding: proxy._config = MagicMock() proxy._config.memory_enabled = False proxy._config.ccr_inject_tool = False + proxy.config = proxy._config proxy._parse_sse_usage_from_buffer = MagicMock(return_value=None) proxy.memory_handler = None return proxy @@ -200,11 +204,18 @@ class TestStreamingRatelimitHeaderForwarding: assert result.headers.get("anthropic-ratelimit-tokens-limit") is None @pytest.mark.asyncio - async def test_upstream_http_error_preserves_status_and_body(self): - """Upstream non-200 streaming responses should preserve HTTP status/body.""" + async def test_upstream_http_error_preserves_status_body_and_metrics(self, monkeypatch): + """Upstream non-200 streaming responses should preserve status/body and metrics.""" proxy = self._create_mock_proxy() mock_response = self._create_mock_upstream_response() mock_response.status_code = 503 + mock_response.headers = httpx.Headers( + { + "content-type": "application/json", + "content-encoding": "gzip", + "content-length": "42", + } + ) mock_response.aread = AsyncMock( return_value=b'{"error":{"message":"capacity exhausted"}}' ) @@ -213,6 +224,8 @@ class TestStreamingRatelimitHeaderForwarding: mock_request = MagicMock() proxy.http_client.build_request = MagicMock(return_value=mock_request) proxy.http_client.send = AsyncMock(return_value=mock_response) + fake_logger = MagicMock() + monkeypatch.setattr(streaming_module, "logger", fake_logger) result = await proxy._stream_response( url="https://api.anthropic.com/v1/messages", @@ -236,6 +249,63 @@ class TestStreamingRatelimitHeaderForwarding: assert result.status_code == 503 assert result.body == b'{"error":{"message":"capacity exhausted"}}' + assert result.headers.get("content-encoding") is None + fake_logger.warning.assert_any_call( + "[%s] Forwarding upstream streaming error status=%s url=%s", + "test-http-error", + 503, + "https://api.anthropic.com/v1/messages", + ) + proxy.metrics.record_request.assert_awaited_once() + proxy.cost_tracker.record_tokens.assert_called_once() + mock_response.aclose.assert_awaited_once() + + @pytest.mark.asyncio + async def test_upstream_http_error_closes_response_when_body_read_fails(self, monkeypatch): + """Reading a streaming error body should still close the upstream response.""" + proxy = self._create_mock_proxy() + mock_response = self._create_mock_upstream_response() + mock_response.status_code = 502 + mock_response.aread = AsyncMock(side_effect=RuntimeError("boom")) + mock_response.aclose = AsyncMock() + + mock_request = MagicMock() + proxy.http_client.build_request = MagicMock(return_value=mock_request) + proxy.http_client.send = AsyncMock(return_value=mock_response) + fake_logger = MagicMock() + monkeypatch.setattr(streaming_module, "logger", fake_logger) + + result = await proxy._stream_response( + url="https://api.anthropic.com/v1/messages", + headers={"x-api-key": "sk-test"}, + body={ + "model": "claude-sonnet-4-20250514", + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + provider="anthropic", + model="claude-sonnet-4-20250514", + request_id="test-http-error-read-fail", + original_tokens=10, + optimized_tokens=10, + tokens_saved=0, + transforms_applied=[], + tags={}, + optimization_latency=0.0, + ) + + assert result.status_code == 502 + assert result.headers.get("content-type") == "application/json" + assert b"Failed to read upstream error response body" in result.body + fake_logger.warning.assert_any_call( + "[%s] Failed reading upstream error body status=%s url=%s error=%s", + "test-http-error-read-fail", + 502, + "https://api.anthropic.com/v1/messages", + mock_response.aread.side_effect, + ) + mock_response.aclose.assert_awaited_once() @pytest.mark.asyncio async def test_connect_error_returns_sse_error(self):