Refactor: split server.py into handler mixins (Steps 5-9)

server.py: 8778 → 2343 lines (73% reduction)

HeadroomProxy now inherits from 5 handler mixins:
- StreamingMixin (901 lines): SSE parsing, streaming, response relay
- AnthropicHandlerMixin (1435 lines): /v1/messages, batch API
- OpenAIHandlerMixin (1320 lines): /v1/chat/completions, /v1/responses, WebSocket
- GeminiHandlerMixin (649 lines): Gemini native API
- BatchHandlerMixin (995 lines): Google/OpenAI batch processing

All backward-compatible imports preserved via re-exports.
181 tests pass, 0 regressions. Real-world tested with Anthropic,
OpenAI, and GPT-5.4 API calls.
This commit is contained in:
chopratejas 2026-04-03 17:08:39 -07:00
parent 6a8ae297d6
commit e8ab444f09
8 changed files with 5341 additions and 5084 deletions

View file

@ -0,0 +1,20 @@
"""Handler mixins for HeadroomProxy.
Each mixin class contains methods extracted from HeadroomProxy that handle
requests for a specific provider or concern. The mixins rely on HeadroomProxy's
__init__ for all self.* attributes (duck typing).
"""
from headroom.proxy.handlers.anthropic import AnthropicHandlerMixin
from headroom.proxy.handlers.batch import BatchHandlerMixin
from headroom.proxy.handlers.gemini import GeminiHandlerMixin
from headroom.proxy.handlers.openai import OpenAIHandlerMixin
from headroom.proxy.handlers.streaming import StreamingMixin
__all__ = [
"AnthropicHandlerMixin",
"BatchHandlerMixin",
"GeminiHandlerMixin",
"OpenAIHandlerMixin",
"StreamingMixin",
]

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,995 @@
"""Batch handler mixin for HeadroomProxy.
Contains all batch API handlers for Google and OpenAI batch operations.
"""
from __future__ import annotations
import json
import logging
import time
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from fastapi import Request
from fastapi.responses import Response
logger = logging.getLogger("headroom.proxy")
class BatchHandlerMixin:
"""Mixin providing batch API handler methods for HeadroomProxy."""
async def handle_google_batch_create(
self,
request: Request,
model: str,
) -> Response:
"""Handle Google POST /v1beta/models/{model}:batchGenerateContent endpoint.
Google batch format:
{
"batch": {
"display_name": "my-batch",
"input_config": {
"requests": {
"requests": [
{
"request": {"contents": [{"parts": [{"text": "..."}]}]},
"metadata": {"key": "request-1"}
}
]
}
}
}
}
This method applies compression to each request's contents before forwarding.
"""
from fastapi.responses import JSONResponse, Response
from headroom.ccr import CCRToolInjector
from headroom.proxy.helpers import MAX_REQUEST_BODY_SIZE, _read_request_json
from headroom.utils import extract_user_query
start_time = time.time()
request_id = await self._next_request_id()
# Check request body size
content_length = request.headers.get("content-length")
if content_length and int(content_length) > MAX_REQUEST_BODY_SIZE:
return JSONResponse(
status_code=413,
content={
"error": {
"code": 413,
"message": f"Request body too large. Maximum size is {MAX_REQUEST_BODY_SIZE // (1024 * 1024)}MB",
"status": "INVALID_ARGUMENT",
}
},
)
# Parse request
try:
body = await _read_request_json(request)
except (json.JSONDecodeError, ValueError) as e:
return JSONResponse(
status_code=400,
content={
"error": {
"code": 400,
"message": f"Invalid request body: {e!s}",
"status": "INVALID_ARGUMENT",
}
},
)
# Extract batch config
batch_config = body.get("batch", {})
input_config = batch_config.get("input_config", {})
requests_wrapper = input_config.get("requests", {})
requests_list = requests_wrapper.get("requests", [])
if not requests_list:
# No inline requests - might be using file input, pass through
logger.debug(f"[{request_id}] Google batch: No inline requests, passing through")
return await self._google_batch_passthrough(request, model, body)
# Extract headers
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
# Track compression stats
total_original_tokens = 0
total_optimized_tokens = 0
total_tokens_saved = 0
compressed_requests = []
pipeline_timing: dict[str, float] = {}
# Apply compression to each request in the batch
for idx, batch_req in enumerate(requests_list):
req_content = batch_req.get("request", {})
metadata = batch_req.get("metadata", {})
contents = req_content.get("contents", [])
if not contents or not self.config.optimize:
# No contents or optimization disabled - pass through unchanged
compressed_requests.append(batch_req)
continue
# Convert Google format to messages for compression
system_instruction = req_content.get("systemInstruction")
messages, preserved_indices = self._gemini_contents_to_messages(
contents, system_instruction
)
# Store original content entries that have non-text parts before compression
preserved_contents = {idx: contents[idx] for idx in preserved_indices}
# Early exit if ALL content has non-text parts (nothing to compress)
if len(preserved_indices) == len(contents):
# All content has non-text parts, skip compression
compressed_requests.append(batch_req)
continue
# Apply optimization
try:
# Default context limit for most models
context_limit = 128000
# Use OpenAI pipeline (similar message format after conversion)
result = self.openai_pipeline.apply(
messages=messages,
model=model,
model_limit=context_limit,
context=extract_user_query(messages),
)
optimized_messages = result.messages
for k, v in result.timing.items():
pipeline_timing[k] = pipeline_timing.get(k, 0.0) + v
# Use pipeline's token counts for consistency with pipeline logs
original_tokens = result.tokens_before
optimized_tokens = result.tokens_after
total_original_tokens += original_tokens
total_optimized_tokens += optimized_tokens
tokens_saved = max(0, original_tokens - optimized_tokens)
total_tokens_saved += tokens_saved
# CCR Tool Injection: Inject retrieval tool if compression occurred
tools = req_content.get("tools")
# Extract existing function declarations if present
existing_funcs = None
if tools:
for tool in tools:
if "functionDeclarations" in tool:
existing_funcs = tool["functionDeclarations"]
break
if self.config.ccr_inject_tool and tokens_saved > 0:
injector = CCRToolInjector(
provider="google",
inject_tool=True,
inject_system_instructions=self.config.ccr_inject_system_instructions,
)
optimized_messages, injected_funcs, was_injected = injector.process_request(
optimized_messages, existing_funcs
)
if was_injected:
logger.debug(
f"[{request_id}] CCR: Injected retrieval tool for Google batch request {idx}"
)
existing_funcs = injected_funcs
# Convert back to Google contents format
optimized_contents, optimized_sys_inst = self._messages_to_gemini_contents(
optimized_messages
)
# Restore preserved content entries that had non-text parts
for orig_idx, original_content in preserved_contents.items():
if orig_idx < len(optimized_contents):
optimized_contents[orig_idx] = original_content
# Create compressed batch request
compressed_req_content = {**req_content, "contents": optimized_contents}
if optimized_sys_inst:
compressed_req_content["systemInstruction"] = optimized_sys_inst
if existing_funcs is not None:
compressed_req_content["tools"] = [{"functionDeclarations": existing_funcs}]
compressed_req = {
"request": compressed_req_content,
"metadata": metadata,
}
compressed_requests.append(compressed_req)
if tokens_saved > 0:
logger.debug(
f"[{request_id}] Google batch request {idx}: "
f"{original_tokens:,} -> {optimized_tokens:,} tokens "
f"(saved {tokens_saved:,})"
)
except Exception as e:
logger.warning(
f"[{request_id}] Optimization failed for Google batch request {idx}: {e}"
)
# Pass through unchanged on failure
compressed_requests.append(batch_req)
total_optimized_tokens += original_tokens
# Update body with compressed requests
body["batch"]["input_config"]["requests"]["requests"] = compressed_requests
optimization_latency = (time.time() - start_time) * 1000
# Forward request to Google
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:batchGenerateContent"
# Add API key to URL if present in headers
api_key = headers.pop("x-goog-api-key", None)
if api_key:
url = f"{url}?key={api_key}"
try:
response = await self._retry_request("POST", url, headers, body)
# Record metrics
await self.metrics.record_request(
provider="google",
model=f"batch:{model}",
input_tokens=total_optimized_tokens,
output_tokens=0,
tokens_saved=total_tokens_saved,
latency_ms=optimization_latency,
overhead_ms=optimization_latency,
pipeline_timing=pipeline_timing,
)
# Log compression stats
if total_tokens_saved > 0:
savings_percent = (
(total_tokens_saved / total_original_tokens * 100)
if total_original_tokens > 0
else 0
)
logger.info(
f"[{request_id}] Google batch compression: "
f"{total_original_tokens:,} -> {total_optimized_tokens:,} tokens "
f"({savings_percent:.1f}% saved across {len(requests_list)} requests)"
)
# Store batch context for CCR result processing
if response.status_code == 200 and self.config.ccr_inject_tool:
try:
response_data = response.json()
batch_name = response_data.get("name")
if batch_name:
await self._store_google_batch_context(
batch_name,
requests_list,
model,
api_key,
)
except Exception as e:
logger.warning(f"[{request_id}] Failed to store Google batch context: {e}")
# Remove compression headers
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
except Exception as e:
logger.error(f"[{request_id}] Google batch request failed: {e}")
return JSONResponse(
status_code=500,
content={
"error": {
"code": 500,
"message": f"Failed to forward batch request: {e!s}",
"status": "INTERNAL",
}
},
)
async def _google_batch_passthrough(
self,
request: Request,
model: str,
body: dict | None = None,
) -> Response:
"""Pass through Google batch request without modification."""
from fastapi.responses import Response
start_time = time.time()
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:batchGenerateContent"
# Add API key to URL if present in headers
api_key = headers.pop("x-goog-api-key", None)
if api_key:
url = f"{url}?key={api_key}"
if body is None:
body_content = await request.body()
else:
body_content = json.dumps(body).encode()
response = await self.http_client.post( # type: ignore[union-attr]
url,
headers=headers,
content=body_content,
)
# Track metrics
latency_ms = (time.time() - start_time) * 1000
await self.metrics.record_request(
provider="google",
model=f"passthrough:batch:{model}",
input_tokens=0,
output_tokens=0,
tokens_saved=0,
latency_ms=latency_ms,
)
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
async def handle_google_batch_passthrough(
self,
request: Request,
batch_name: str | None = None,
) -> Response:
"""Handle Google batch passthrough endpoints.
Used for:
- GET /v1beta/batches/{batch_name} - Get batch status
- POST /v1beta/batches/{batch_name}:cancel - Cancel batch
- DELETE /v1beta/batches/{batch_name} - Delete batch
"""
from fastapi.responses import Response
start_time = time.time()
path = request.url.path
url = f"{self.GEMINI_API_URL}{path}"
# Preserve query string parameters
if request.url.query:
url = f"{url}?{request.url.query}"
headers = dict(request.headers.items())
headers.pop("host", None)
# Handle API key
api_key = headers.pop("x-goog-api-key", None)
if api_key:
if "?" in url:
url = f"{url}&key={api_key}"
else:
url = f"{url}?key={api_key}"
body = await request.body()
response = await self.http_client.request( # type: ignore[union-attr]
method=request.method,
url=url,
headers=headers,
content=body,
)
# Track metrics
latency_ms = (time.time() - start_time) * 1000
await self.metrics.record_request(
provider="google",
model="passthrough:batches",
input_tokens=0,
output_tokens=0,
tokens_saved=0,
latency_ms=latency_ms,
)
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
async def _store_google_batch_context(
self,
batch_name: str,
requests_list: list[dict[str, Any]],
model: str,
api_key: str | None,
) -> None:
"""Store Google batch context for CCR result processing.
Args:
batch_name: The batch name from the API response.
requests_list: The original batch requests.
model: The model used for the batch.
api_key: The API key for continuation calls.
"""
from headroom.ccr import BatchContext, BatchRequestContext, get_batch_context_store
store = get_batch_context_store()
context = BatchContext(
batch_id=batch_name,
provider="google",
api_key=api_key,
api_base_url=self.GEMINI_API_URL,
)
for batch_req in requests_list:
metadata = batch_req.get("metadata", {})
custom_id = metadata.get("key", "")
req_content = batch_req.get("request", {})
contents = req_content.get("contents", [])
system_instruction = req_content.get("systemInstruction")
# Convert contents to messages format for CCR handler
messages, _ = self._gemini_contents_to_messages(contents, system_instruction)
# Extract system instruction text if present
sys_text = None
if system_instruction:
parts = system_instruction.get("parts", [])
if parts and isinstance(parts[0], dict):
sys_text = parts[0].get("text")
context.add_request(
BatchRequestContext(
custom_id=custom_id,
messages=messages,
tools=req_content.get("tools"),
model=model,
system_instruction=sys_text,
)
)
await store.store(context)
logger.debug(
f"Stored Google batch context for {batch_name} with {len(requests_list)} requests"
)
async def handle_google_batch_results(
self,
request: Request,
batch_name: str,
) -> Response:
"""Handle Google batch results with CCR post-processing.
Google batch results endpoint returns the batch operation status.
When status is SUCCEEDED, results are embedded in the response.
This handler processes CCR tool calls in those results.
"""
from fastapi.responses import JSONResponse, Response
from headroom.ccr import BatchResultProcessor, get_batch_context_store
start_time = time.time()
# Forward request to get batch status/results
url = f"{self.GEMINI_API_URL}/v1beta/{batch_name}"
if request.url.query:
url = f"{url}?{request.url.query}"
headers = dict(request.headers.items())
headers.pop("host", None)
# Handle API key
api_key = headers.pop("x-goog-api-key", None)
if api_key:
if "?" in url:
url = f"{url}&key={api_key}"
else:
url = f"{url}?key={api_key}"
response = await self.http_client.get(url, headers=headers) # type: ignore[union-attr]
if response.status_code != 200:
# Error - pass through
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Parse response
try:
response_data = response.json()
except json.JSONDecodeError:
# Not JSON - pass through
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Check if batch has results (state must be SUCCEEDED)
metadata = response_data.get("metadata", {})
state = metadata.get("state")
if state != "SUCCEEDED":
# Batch not complete - pass through
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Extract results from response
# Google embeds results in the batch response
results = response_data.get("response", {}).get("responses", [])
if not results:
# No results to process
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Check if we have context and CCR processing is enabled
store = get_batch_context_store()
batch_context = await store.get(batch_name)
if batch_context is None or not self.config.ccr_inject_tool:
# No context or CCR disabled - pass through
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Process results with CCR handler
processor = BatchResultProcessor(self.http_client) # type: ignore[arg-type]
processed = await processor.process_results(batch_name, results, "google")
# Update response with processed results
processed_results = [p.result for p in processed]
response_data["response"]["responses"] = processed_results
for p in processed:
if p.was_processed:
logger.info(
f"CCR: Processed Google batch result {p.custom_id} "
f"({p.continuation_rounds} continuation rounds)"
)
# Track metrics
latency_ms = (time.time() - start_time) * 1000
await self.metrics.record_request(
provider="google",
model="batch:ccr-processed",
input_tokens=0,
output_tokens=0,
tokens_saved=0,
latency_ms=latency_ms,
)
return JSONResponse(content=response_data, status_code=200)
async def handle_batch_create(self, request: Request) -> Response:
"""Handle POST /v1/batches - Create a batch with compression.
Flow:
1. Parse request to get input_file_id
2. Download the JSONL file content from OpenAI
3. Parse each line and compress the messages
4. Create a new compressed JSONL file
5. Upload compressed file to OpenAI
6. Create batch with the new compressed file_id
7. Return batch object with compression stats in metadata
"""
from fastapi.responses import JSONResponse, Response
from headroom.proxy.helpers import _read_request_json
start_time = time.time()
request_id = await self._next_request_id()
# Parse request
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}",
"type": "invalid_request_error",
"code": "invalid_json",
}
},
)
input_file_id = body.get("input_file_id")
endpoint = body.get("endpoint")
completion_window = body.get("completion_window", "24h")
metadata = body.get("metadata", {})
if not input_file_id:
return JSONResponse(
status_code=400,
content={
"error": {
"message": "input_file_id is required",
"type": "invalid_request_error",
"code": "missing_parameter",
}
},
)
if not endpoint:
return JSONResponse(
status_code=400,
content={
"error": {
"message": "endpoint is required",
"type": "invalid_request_error",
"code": "missing_parameter",
}
},
)
# Only compress chat completions endpoint
if endpoint != "/v1/chat/completions":
# Pass through for other endpoints
return await self._batch_passthrough(request, body)
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
try:
# Step 1: Download the input file from OpenAI
logger.info(f"[{request_id}] Batch: Downloading input file {input_file_id}")
file_content = await self._download_openai_file(input_file_id, headers)
if file_content is None:
return JSONResponse(
status_code=404,
content={
"error": {
"message": f"Failed to download file {input_file_id}",
"type": "invalid_request_error",
"code": "file_not_found",
}
},
)
# Step 2: Parse and compress each line
logger.info(f"[{request_id}] Batch: Compressing JSONL content")
compressed_lines, stats = await self._compress_batch_jsonl(file_content, request_id)
if stats["total_requests"] == 0:
return JSONResponse(
status_code=400,
content={
"error": {
"message": "No valid requests found in input file",
"type": "invalid_request_error",
"code": "empty_file",
}
},
)
# Step 3: Create compressed JSONL content
compressed_content = "\n".join(compressed_lines)
# Step 4: Upload compressed file to OpenAI
logger.info(f"[{request_id}] Batch: Uploading compressed file")
new_file_id = await self._upload_openai_file(
compressed_content, f"compressed_{input_file_id}.jsonl", headers
)
if new_file_id is None:
return JSONResponse(
status_code=500,
content={
"error": {
"message": "Failed to upload compressed file",
"type": "server_error",
"code": "upload_failed",
}
},
)
# Step 5: Create batch with compressed file
logger.info(f"[{request_id}] Batch: Creating batch with compressed file {new_file_id}")
# Add compression stats to metadata
compression_metadata = {
**metadata,
"headroom_compressed": "true",
"headroom_original_file_id": input_file_id,
"headroom_total_requests": str(stats["total_requests"]),
"headroom_tokens_saved": str(stats["total_tokens_saved"]),
"headroom_original_tokens": str(stats["total_original_tokens"]),
"headroom_compressed_tokens": str(stats["total_compressed_tokens"]),
"headroom_savings_percent": f"{stats['savings_percent']:.1f}",
}
batch_body = {
"input_file_id": new_file_id,
"endpoint": endpoint,
"completion_window": completion_window,
"metadata": compression_metadata,
}
url = f"{self.OPENAI_API_URL}/v1/batches"
response = await self.http_client.post(url, json=batch_body, headers=headers) # type: ignore[union-attr]
total_latency = (time.time() - start_time) * 1000
# Log compression stats
logger.info(
f"[{request_id}] Batch created: {stats['total_requests']} requests, "
f"{stats['total_original_tokens']:,} -> {stats['total_compressed_tokens']:,} tokens "
f"(saved {stats['total_tokens_saved']:,} tokens, {stats['savings_percent']:.1f}%) "
f"in {total_latency:.0f}ms"
)
# Record metrics
await self.metrics.record_request(
provider="openai",
model="batch",
input_tokens=stats["total_compressed_tokens"],
output_tokens=0,
tokens_saved=stats["total_tokens_saved"],
latency_ms=total_latency,
)
# Return response with compression info in headers
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
response_headers["x-headroom-tokens-saved"] = str(stats["total_tokens_saved"])
response_headers["x-headroom-savings-percent"] = f"{stats['savings_percent']:.1f}"
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
except Exception as e:
logger.error(f"[{request_id}] Batch creation failed: {type(e).__name__}: {e}")
await self.metrics.record_failed()
return JSONResponse(
status_code=500,
content={
"error": {
"message": "An error occurred while processing the batch request",
"type": "server_error",
"code": "batch_processing_error",
}
},
)
async def _download_openai_file(self, file_id: str, headers: dict) -> str | None:
"""Download file content from OpenAI."""
url = f"{self.OPENAI_API_URL}/v1/files/{file_id}/content"
try:
response = await self.http_client.get(url, headers=headers) # type: ignore[union-attr]
if response.status_code == 200:
return str(response.text)
logger.error(f"Failed to download file {file_id}: {response.status_code}")
return None
except Exception as e:
logger.error(f"Error downloading file {file_id}: {e}")
return None
async def _upload_openai_file(self, content: str, filename: str, headers: dict) -> str | None:
"""Upload a file to OpenAI for batch processing."""
url = f"{self.OPENAI_API_URL}/v1/files"
# Prepare multipart form data
# We need to use httpx's files parameter for multipart upload
files = {
"file": (filename, content.encode("utf-8"), "application/jsonl"),
}
data = {
"purpose": "batch",
}
# Remove content-type from headers (httpx will set it for multipart)
upload_headers = {k: v for k, v in headers.items() if k.lower() != "content-type"}
try:
response = await self.http_client.post( # type: ignore[union-attr]
url, files=files, data=data, headers=upload_headers
)
if response.status_code == 200:
result = response.json()
file_id: str | None = result.get("id")
return file_id
logger.error(f"Failed to upload file: {response.status_code} - {response.text}")
return None
except Exception as e:
logger.error(f"Error uploading file: {e}")
return None
async def _compress_batch_jsonl(self, content: str, request_id: str) -> tuple[list[str], dict]:
"""Compress messages in each line of a batch JSONL file.
Returns:
Tuple of (compressed_lines, stats_dict)
"""
from headroom.ccr import CCRToolInjector
from headroom.tokenizers import get_tokenizer
from headroom.utils import extract_user_query
lines = content.strip().split("\n")
compressed_lines = []
total_original_tokens = 0
total_compressed_tokens = 0
total_requests = 0
errors = 0
tokenizer = get_tokenizer("gpt-4") # Use gpt-4 tokenizer for batch
for i, line in enumerate(lines):
if not line.strip():
continue
try:
request_obj = json.loads(line)
body = request_obj.get("body", {})
messages = body.get("messages", [])
model = body.get("model", "gpt-4")
if not messages:
# No messages to compress, pass through
compressed_lines.append(line)
total_requests += 1
continue
# Compress messages using the OpenAI pipeline
if self.config.optimize:
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),
)
compressed_messages = result.messages
# Use pipeline's token counts for consistency with pipeline logs
original_tokens = result.tokens_before
compressed_tokens = result.tokens_after
except Exception as e:
logger.warning(f"[{request_id}] Compression failed for line {i}: {e}")
compressed_messages = messages
original_tokens = tokenizer.count_messages(messages)
compressed_tokens = original_tokens
else:
compressed_messages = messages
original_tokens = tokenizer.count_messages(messages)
compressed_tokens = original_tokens
total_original_tokens += original_tokens
total_compressed_tokens += compressed_tokens
tokens_saved = original_tokens - compressed_tokens
# CCR Tool Injection: Inject retrieval tool if compression occurred
tools = body.get("tools")
if self.config.ccr_inject_tool and tokens_saved > 0:
injector = CCRToolInjector(
provider="openai",
inject_tool=True,
inject_system_instructions=self.config.ccr_inject_system_instructions,
)
compressed_messages, tools, was_injected = injector.process_request(
compressed_messages, tools
)
if was_injected:
logger.debug(
f"[{request_id}] CCR: Injected retrieval tool for batch line {i}"
)
# Update body with compressed messages
body["messages"] = compressed_messages
if tools is not None:
body["tools"] = tools
request_obj["body"] = body
compressed_lines.append(json.dumps(request_obj))
total_requests += 1
except json.JSONDecodeError as e:
logger.warning(f"[{request_id}] Invalid JSON on line {i}: {e}")
errors += 1
# Keep original line on error
compressed_lines.append(line)
total_requests += 1
total_tokens_saved = total_original_tokens - total_compressed_tokens
savings_percent = (
(total_tokens_saved / total_original_tokens * 100) if total_original_tokens > 0 else 0
)
stats = {
"total_requests": total_requests,
"total_original_tokens": total_original_tokens,
"total_compressed_tokens": total_compressed_tokens,
"total_tokens_saved": total_tokens_saved,
"savings_percent": savings_percent,
"errors": errors,
}
return compressed_lines, stats
async def _batch_passthrough(self, request: Request, body: dict) -> Response:
"""Pass through batch request to OpenAI without compression."""
from fastapi.responses import Response
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
url = f"{self.OPENAI_API_URL}/v1/batches"
response = await self.http_client.post(url, json=body, headers=headers) # type: ignore[union-attr]
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
async def handle_batch_list(self, request: Request) -> Response:
"""Handle GET /v1/batches - List batches (passthrough)."""
return await self.handle_passthrough(request, self.OPENAI_API_URL)
async def handle_batch_get(self, request: Request, batch_id: str) -> Response:
"""Handle GET /v1/batches/{batch_id} - Get batch (passthrough)."""
return await self.handle_passthrough(request, self.OPENAI_API_URL)
async def handle_batch_cancel(self, request: Request, batch_id: str) -> Response:
"""Handle POST /v1/batches/{batch_id}/cancel - Cancel batch (passthrough)."""
return await self.handle_passthrough(request, self.OPENAI_API_URL)

View file

@ -0,0 +1,649 @@
"""Gemini handler mixin for HeadroomProxy.
Contains all Google Gemini API handlers including format conversion utilities.
"""
from __future__ import annotations
import json
import logging
import time
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastapi import Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
logger = logging.getLogger("headroom.proxy")
class GeminiHandlerMixin:
"""Mixin providing Gemini API handler methods for HeadroomProxy."""
def _has_non_text_parts(self, content: dict) -> bool:
"""Check if a Gemini content entry has non-text parts.
Non-text parts include:
- inlineData: Base64-encoded images/media
- fileData: File references (URI + MIME type)
- functionCall: Function calls from model
- functionResponse: Responses to function calls
Args:
content: A single Gemini content entry with 'parts' list.
Returns:
True if any part contains non-text data.
"""
parts = content.get("parts", [])
for part in parts:
if any(
key in part
for key in ("inlineData", "fileData", "functionCall", "functionResponse")
):
return True
return False
def _gemini_contents_to_messages(
self, contents: list[dict], system_instruction: dict | None = None
) -> tuple[list[dict], set[int]]:
"""Convert Gemini contents[] format to OpenAI messages[] format for optimization.
Gemini format:
contents: [{"role": "user", "parts": [{"text": "..."}]}]
systemInstruction: {"parts": [{"text": "..."}]}
OpenAI format:
messages: [{"role": "user", "content": "..."}]
Returns:
Tuple of (messages, preserved_indices) where preserved_indices contains
the indices of content entries that have non-text parts (images, function
calls, etc.) and should not be compressed.
"""
messages = []
preserved_indices: set[int] = set()
# Add system instruction as system message
if system_instruction:
parts = system_instruction.get("parts", [])
text_parts = [p.get("text", "") for p in parts if "text" in p]
if text_parts:
messages.append({"role": "system", "content": "\n".join(text_parts)})
# Convert contents to messages
for idx, content in enumerate(contents):
# Track content entries with non-text parts
if self._has_non_text_parts(content):
preserved_indices.add(idx)
role = content.get("role", "user")
# Map Gemini roles to OpenAI roles
if role == "model":
role = "assistant"
parts = content.get("parts", [])
text_parts = [p.get("text", "") for p in parts if "text" in p]
if text_parts:
messages.append({"role": role, "content": "\n".join(text_parts)})
return messages, preserved_indices
def _messages_to_gemini_contents(self, messages: list[dict]) -> tuple[list[dict], dict | None]:
"""Convert OpenAI messages[] format back to Gemini contents[] format.
Returns:
(contents, system_instruction) tuple
"""
contents = []
system_instruction = None
for msg in messages:
role = msg.get("role", "user")
content = msg.get("content", "")
if role == "system":
# Extract as systemInstruction
system_instruction = {"parts": [{"text": content}]}
else:
# Map OpenAI roles to Gemini roles
gemini_role = "model" if role == "assistant" else "user"
contents.append({"role": gemini_role, "parts": [{"text": content}]})
return contents, system_instruction
async def handle_gemini_generate_content(
self,
request: Request,
model: str,
) -> Response | StreamingResponse:
"""Handle Gemini native /v1beta/models/{model}:generateContent endpoint.
Gemini's native API differs from OpenAI:
- Input: `contents[]` with `parts[]` instead of `messages`
- System: `systemInstruction` instead of system message
- Auth: `x-goog-api-key` header instead of `Authorization: Bearer`
- Output: `candidates[].content.parts[].text`
"""
from fastapi import HTTPException
from fastapi.responses import JSONResponse, Response
from headroom.proxy.helpers import MAX_REQUEST_BODY_SIZE, _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()
# Check request body size
content_length = request.headers.get("content-length")
if content_length and int(content_length) > MAX_REQUEST_BODY_SIZE:
return JSONResponse(
status_code=413,
content={
"error": {
"message": f"Request body too large. Maximum size is {MAX_REQUEST_BODY_SIZE // (1024 * 1024)}MB",
"code": 413,
}
},
)
# Parse request
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,
}
},
)
contents = body.get("contents", [])
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
tags = self._extract_tags(headers)
# Rate limiting (use Gemini API key)
if self.rate_limiter:
rate_key = headers.get("x-goog-api-key", "default")[:20]
allowed, wait_seconds = await self.rate_limiter.check_request(rate_key)
if not allowed:
await self.metrics.record_rate_limited()
raise HTTPException(
status_code=429,
detail=f"Rate limited. Retry after {wait_seconds:.1f}s",
)
# Convert Gemini format to messages for optimization
system_instruction = body.get("systemInstruction")
messages, preserved_indices = self._gemini_contents_to_messages(
contents, system_instruction
)
# Store original content entries that have non-text parts before compression
preserved_contents = {idx: contents[idx] for idx in preserved_indices}
# Early exit if ALL content has non-text parts (nothing to compress)
if len(preserved_indices) == len(contents):
# All content has non-text parts, skip compression entirely
# Just forward the request as-is
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:generateContent"
query_params = dict(request.query_params)
is_streaming = query_params.get("alt") == "sse"
if "key" in query_params:
url += f"?key={query_params['key']}"
if is_streaming:
stream_url = (
f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?alt=sse"
)
if "key" in query_params:
stream_url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?key={query_params['key']}&alt=sse"
return await self._stream_response(
stream_url,
headers,
body,
"gemini",
model,
request_id,
0,
0,
0,
[],
tags,
0,
)
else:
response = await self._retry_request("POST", url, headers, body)
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Token counting
tokenizer = get_tokenizer(model)
original_tokens = tokenizer.count_messages(messages)
# Optimization
transforms_applied: list[str] = []
waste_signals_dict: dict[str, int] | None = None
optimized_messages = messages
optimized_tokens = original_tokens
_compression_failed = False
_license_ok = self.usage_reporter.should_compress if self.usage_reporter else True
if self.config.optimize and messages and _license_ok:
try:
# Use OpenAI pipeline (similar message format)
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
# Use pipeline's token counts for consistency with pipeline logs
original_tokens = result.tokens_before
optimized_tokens = result.tokens_after
if result.waste_signals:
waste_signals_dict = result.waste_signals.to_dict()
except Exception as e:
_compression_failed = True
logger.warning(f"[{request_id}] Gemini optimization failed: {e}")
tokens_saved = max(0, original_tokens - optimized_tokens)
optimization_latency = (time.time() - start_time) * 1000
# Query Echo: disabled — hurts prefix caching in long conversations.
# Convert back to Gemini format if optimized
if optimized_messages != messages:
optimized_contents, optimized_system = self._messages_to_gemini_contents(
optimized_messages
)
# Restore preserved content entries that had non-text parts
for orig_idx, original_content in preserved_contents.items():
if orig_idx < len(optimized_contents):
optimized_contents[orig_idx] = original_content
body["contents"] = optimized_contents
if optimized_system:
body["systemInstruction"] = optimized_system
elif "systemInstruction" in body:
del body["systemInstruction"]
# Build URL - model is extracted from path
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:generateContent"
# Check if streaming requested via query param
query_params = dict(request.query_params)
is_streaming = query_params.get("alt") == "sse"
# Preserve API key in query params if present
if "key" in query_params:
url += f"?key={query_params['key']}"
try:
if is_streaming:
# For streaming, use streamGenerateContent endpoint
stream_url = (
f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?alt=sse"
)
if "key" in query_params:
stream_url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?key={query_params['key']}&alt=sse"
return await self._stream_response(
stream_url,
headers,
body,
"gemini",
model,
request_id,
original_tokens,
optimized_tokens,
tokens_saved,
transforms_applied,
tags,
optimization_latency,
)
else:
response = await self._retry_request("POST", url, headers, body)
total_latency = (time.time() - start_time) * 1000
total_input_tokens = optimized_tokens # fallback
output_tokens = 0
cache_read_tokens = 0
try:
resp_json = response.json()
usage = resp_json.get("usageMetadata", {})
total_input_tokens = usage.get("promptTokenCount", optimized_tokens)
output_tokens = usage.get("candidatesTokenCount", 0)
# Gemini returns cachedContentTokenCount for context-cached tokens
# These are charged at 10-25% of the input price depending on model
cache_read_tokens = usage.get("cachedContentTokenCount", 0)
except (KeyError, TypeError, AttributeError) as e:
logger.debug(
f"[{request_id}] Failed to extract cached tokens from Gemini response: {e}"
)
if self.cost_tracker:
self.cost_tracker.record_tokens(
model,
tokens_saved,
optimized_tokens,
cache_read_tokens=cache_read_tokens,
)
await self.metrics.record_request(
provider="gemini",
model=model,
input_tokens=total_input_tokens,
output_tokens=output_tokens,
tokens_saved=tokens_saved,
latency_ms=total_latency,
overhead_ms=optimization_latency,
waste_signals=waste_signals_dict,
cache_read_tokens=cache_read_tokens,
)
if tokens_saved > 0:
logger.info(
f"[{request_id}] Gemini {model}: {original_tokens:,}{optimized_tokens:,} "
f"(saved {tokens_saved:,} tokens)"
)
else:
logger.info(f"[{request_id}] Gemini {model}: {original_tokens:,} tokens")
# Remove compression headers
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
# Inject Headroom compression metrics (for SaaS metering)
response_headers["x-headroom-tokens-before"] = str(original_tokens)
response_headers["x-headroom-tokens-after"] = str(optimized_tokens)
response_headers["x-headroom-tokens-saved"] = str(tokens_saved)
response_headers["x-headroom-model"] = model
if transforms_applied:
response_headers["x-headroom-transforms"] = ",".join(transforms_applied)
if cache_read_tokens > 0:
response_headers["x-headroom-cached"] = "true"
if _compression_failed:
response_headers["x-headroom-compression-failed"] = "true"
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
except Exception as e:
await self.metrics.record_failed()
logger.error(f"[{request_id}] Gemini request failed: {type(e).__name__}: {e}")
return JSONResponse(
status_code=502,
content={
"error": {
"message": "An error occurred while processing your request. Please try again.",
"code": 502,
}
},
)
async def handle_gemini_stream_generate_content(
self,
request: Request,
model: str,
) -> StreamingResponse | JSONResponse:
"""Handle Gemini streaming endpoint /v1beta/models/{model}:streamGenerateContent."""
from fastapi.responses import JSONResponse
from headroom.proxy.helpers import _read_request_json
from headroom.tokenizers import get_tokenizer
start_time = time.time()
request_id = await self._next_request_id()
# Parse request
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,
}
},
)
contents = body.get("contents", [])
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
tags = self._extract_tags(headers)
# Token counting
tokenizer = get_tokenizer(model)
original_tokens = 0
for content in contents:
parts = content.get("parts", [])
for part in parts:
if "text" in part:
original_tokens += tokenizer.count_text(part["text"])
optimization_latency = (time.time() - start_time) * 1000
# Build URL with SSE param
query_params = dict(request.query_params)
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?alt=sse"
if "key" in query_params:
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:streamGenerateContent?key={query_params['key']}&alt=sse"
return await self._stream_response(
url,
headers,
body,
"gemini",
model,
request_id,
original_tokens,
original_tokens,
0, # tokens_saved
[], # transforms_applied
tags,
optimization_latency,
)
async def handle_gemini_count_tokens(
self,
request: Request,
model: str,
) -> Response:
"""Handle Gemini /v1beta/models/{model}:countTokens endpoint with compression.
This endpoint counts tokens AFTER applying compression, so users can see
how many tokens they'll actually use after optimization.
The request format is the same as generateContent:
{"contents": [...], "systemInstruction": {...}}
"""
from fastapi.responses import JSONResponse, Response
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()
# Parse request
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,
}
},
)
contents = body.get("contents", [])
headers = dict(request.headers.items())
headers.pop("host", None)
headers.pop("content-length", None)
# Convert Gemini format to messages for optimization
system_instruction = body.get("systemInstruction")
messages, preserved_indices = self._gemini_contents_to_messages(
contents, system_instruction
)
# Store original content entries that have non-text parts before compression
preserved_contents = {idx: contents[idx] for idx in preserved_indices}
# Early exit if ALL content has non-text parts (nothing to compress)
if len(preserved_indices) == len(contents):
# All content has non-text parts, skip compression entirely
# Just forward the countTokens request as-is
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:countTokens"
query_params = dict(request.query_params)
if "key" in query_params:
url += f"?key={query_params['key']}"
response = await self._retry_request("POST", url, headers, body)
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
# Token counting (original)
tokenizer = get_tokenizer(model)
original_tokens = tokenizer.count_messages(messages)
# Apply compression using the same pipeline as generateContent
transforms_applied: list[str] = []
optimized_messages = messages
if self.config.optimize and messages:
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
except Exception as e:
logger.warning(f"[{request_id}] Gemini countTokens optimization failed: {e}")
# Convert back to Gemini format for the API call
if optimized_messages != messages:
optimized_contents, optimized_system = self._messages_to_gemini_contents(
optimized_messages
)
# Restore preserved content entries that had non-text parts
for orig_idx, original_content in preserved_contents.items():
if orig_idx < len(optimized_contents):
optimized_contents[orig_idx] = original_content
body["contents"] = optimized_contents
if optimized_system:
body["systemInstruction"] = optimized_system
elif "systemInstruction" in body:
del body["systemInstruction"]
# Build URL
url = f"{self.GEMINI_API_URL}/v1beta/models/{model}:countTokens"
# Preserve API key in query params if present
query_params = dict(request.query_params)
if "key" in query_params:
url += f"?key={query_params['key']}"
try:
response = await self._retry_request("POST", url, headers, body)
total_latency = (time.time() - start_time) * 1000
# Parse response to get token count
compressed_tokens = 0
try:
resp_json = response.json()
compressed_tokens = resp_json.get("totalTokens", 0)
except (json.JSONDecodeError, ValueError) as e:
logger.debug(f"[{request_id}] Failed to parse Gemini token count response: {e}")
# Track stats
tokens_saved = (
max(0, original_tokens - compressed_tokens) if compressed_tokens > 0 else 0
)
await self.metrics.record_request(
provider="gemini",
model=model,
input_tokens=compressed_tokens,
output_tokens=0,
tokens_saved=tokens_saved,
latency_ms=total_latency,
)
if tokens_saved > 0:
logger.info(
f"[{request_id}] Gemini countTokens {model}: {original_tokens:,}{compressed_tokens:,} "
f"(saved {tokens_saved:,} tokens, transforms: {transforms_applied})"
)
else:
logger.info(
f"[{request_id}] Gemini countTokens {model}: {compressed_tokens:,} tokens"
)
# Remove compression headers
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return Response(
content=response.content,
status_code=response.status_code,
headers=response_headers,
)
except Exception as e:
await self.metrics.record_failed()
logger.error(f"[{request_id}] Gemini countTokens failed: {type(e).__name__}: {e}")
return JSONResponse(
status_code=502,
content={
"error": {
"message": "An error occurred while processing your request. Please try again.",
"code": 502,
}
},
)

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,901 @@
"""Streaming handler mixin for HeadroomProxy.
Contains SSE parsing, streaming response generation, and related utilities.
"""
from __future__ import annotations
import contextlib
import json
import logging
import time
from datetime import datetime
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from fastapi.responses import StreamingResponse
import httpx
logger = logging.getLogger("headroom.proxy")
class StreamingMixin:
"""Mixin providing streaming response methods for HeadroomProxy."""
def _parse_sse_usage(self, chunk: bytes, provider: str) -> dict[str, int] | None:
"""Parse usage information from SSE chunk.
For Anthropic: Looks for message_start (input tokens) and message_delta (output tokens)
For OpenAI: Looks for final chunk with usage object (requires stream_options.include_usage=true)
For Gemini: Looks for usageMetadata in each chunk
Returns dict with keys: input_tokens, output_tokens, cache_read_input_tokens, cache_creation_input_tokens
Returns None if no usage found in this chunk.
"""
try:
text = chunk.decode("utf-8", errors="ignore")
# SSE format: "data: {...}\n\n" or "event: ...\ndata: {...}\n\n"
for line in text.split("\n"):
if not line.startswith("data: "):
continue
data_str = line[6:].strip()
if not data_str or data_str == "[DONE]":
continue
try:
data = json.loads(data_str)
except json.JSONDecodeError:
continue
usage = {}
if provider == "anthropic":
# Anthropic sends message_start with input tokens
# and message_delta with output tokens
event_type = data.get("type", "")
if event_type == "message_start":
msg = data.get("message", {})
msg_usage = msg.get("usage", {})
if msg_usage:
usage["input_tokens"] = msg_usage.get("input_tokens", 0)
usage["cache_read_input_tokens"] = msg_usage.get(
"cache_read_input_tokens", 0
)
usage["cache_creation_input_tokens"] = msg_usage.get(
"cache_creation_input_tokens", 0
)
elif event_type == "message_delta":
delta_usage = data.get("usage", {})
if delta_usage:
usage["output_tokens"] = delta_usage.get("output_tokens", 0)
elif provider == "openai":
# OpenAI sends usage in final chunk (when stream_options.include_usage=true)
chunk_usage = data.get("usage")
if chunk_usage:
usage["input_tokens"] = chunk_usage.get("prompt_tokens", 0)
usage["output_tokens"] = chunk_usage.get("completion_tokens", 0)
# OpenAI has cached tokens in prompt_tokens_details
details = chunk_usage.get("prompt_tokens_details", {})
usage["cache_read_input_tokens"] = details.get("cached_tokens", 0)
elif provider == "gemini":
# Gemini sends usageMetadata in each streaming chunk
# Format: {"usageMetadata": {"promptTokenCount": N, "candidatesTokenCount": M}}
usage_meta = data.get("usageMetadata")
if usage_meta:
usage["input_tokens"] = usage_meta.get("promptTokenCount", 0)
usage["output_tokens"] = usage_meta.get("candidatesTokenCount", 0)
# Gemini also has cachedContentTokenCount for context caching
usage["cache_read_input_tokens"] = usage_meta.get(
"cachedContentTokenCount", 0
)
if usage:
return usage
except (UnicodeDecodeError, KeyError, TypeError) as e:
# Don't fail streaming on parse errors
logger.debug(f"SSE usage parsing error for {provider}: {e}")
return None
def _parse_sse_usage_from_buffer(
self, stream_state: dict[str, Any], provider: str
) -> dict[str, int] | None:
"""Parse usage from buffered SSE data, handling split chunks.
Processes complete SSE events (ending with double newline) from the buffer
and removes them from the buffer. Incomplete events are kept in the buffer
for the next chunk.
"""
buffer = stream_state["sse_buffer"]
usage_found: dict[str, int] = {}
# Process complete SSE events (separated by double newlines)
while "\n\n" in buffer:
event_end = buffer.index("\n\n")
event_text = buffer[: event_end + 2]
buffer = buffer[event_end + 2 :]
# Parse this complete event
for line in event_text.split("\n"):
if not line.startswith("data: "):
continue
data_str = line[6:].strip()
if not data_str or data_str == "[DONE]":
continue
try:
data = json.loads(data_str)
except json.JSONDecodeError:
continue
if provider == "anthropic":
event_type = data.get("type", "")
if event_type == "message_start":
msg = data.get("message", {})
msg_usage = msg.get("usage", {})
if msg_usage:
usage_found["input_tokens"] = msg_usage.get("input_tokens", 0)
usage_found["cache_read_input_tokens"] = msg_usage.get(
"cache_read_input_tokens", 0
)
usage_found["cache_creation_input_tokens"] = msg_usage.get(
"cache_creation_input_tokens", 0
)
# INFO logging for cache token tracking (temporary for debugging)
logger.info(
f"[CACHE] Anthropic usage: input={usage_found.get('input_tokens')}, "
f"cache_read={usage_found.get('cache_read_input_tokens')}, "
f"cache_write={usage_found.get('cache_creation_input_tokens')}"
)
elif event_type == "message_delta":
delta_usage = data.get("usage", {})
if delta_usage:
usage_found["output_tokens"] = delta_usage.get("output_tokens", 0)
elif provider == "openai":
chunk_usage = data.get("usage")
if chunk_usage:
usage_found["input_tokens"] = chunk_usage.get("prompt_tokens", 0)
usage_found["output_tokens"] = chunk_usage.get("completion_tokens", 0)
details = chunk_usage.get("prompt_tokens_details", {})
usage_found["cache_read_input_tokens"] = details.get("cached_tokens", 0)
elif provider == "gemini":
usage_meta = data.get("usageMetadata")
if usage_meta:
usage_found["input_tokens"] = usage_meta.get("promptTokenCount", 0)
usage_found["output_tokens"] = usage_meta.get("candidatesTokenCount", 0)
usage_found["cache_read_input_tokens"] = usage_meta.get(
"cachedContentTokenCount", 0
)
# Update buffer with remaining incomplete data
stream_state["sse_buffer"] = buffer
return usage_found if usage_found else None
def _parse_sse_to_response(self, sse_data: str, provider: str) -> dict[str, Any] | None:
"""Parse SSE data to reconstruct the API response JSON.
Args:
sse_data: Raw SSE data string.
provider: Provider type for parsing.
Returns:
Reconstructed response dict or None if parsing fails.
"""
if provider != "anthropic":
return None # Only implemented for Anthropic
response: dict[str, Any] = {"content": [], "usage": {}}
current_block: dict[str, Any] | None = None
for line in sse_data.split("\n"):
if not line.startswith("data: "):
continue
data_str = line[6:].strip()
if not data_str or data_str == "[DONE]":
continue
try:
data = json.loads(data_str)
except json.JSONDecodeError:
continue
event_type = data.get("type", "")
if event_type == "message_start":
msg = data.get("message", {})
response["id"] = msg.get("id")
response["model"] = msg.get("model")
response["role"] = msg.get("role", "assistant")
response["stop_reason"] = msg.get("stop_reason")
if msg.get("usage"):
response["usage"].update(msg["usage"])
elif event_type == "content_block_start":
block = data.get("content_block", {})
current_block = {
"type": block.get("type"),
"index": data.get("index", len(response["content"])),
}
if block.get("type") == "text":
current_block["text"] = block.get("text", "")
elif block.get("type") == "tool_use":
current_block["id"] = block.get("id")
current_block["name"] = block.get("name")
current_block["input"] = {}
elif event_type == "content_block_delta":
if current_block:
delta = data.get("delta", {})
if delta.get("type") == "text_delta":
current_block["text"] = current_block.get("text", "") + delta.get(
"text", ""
)
elif delta.get("type") == "input_json_delta":
# Accumulate partial JSON for tool input
partial = delta.get("partial_json", "")
current_block["_partial_json"] = (
current_block.get("_partial_json", "") + partial
)
elif event_type == "content_block_stop":
if current_block:
# Parse accumulated JSON for tool_use blocks
if current_block.get("type") == "tool_use" and "_partial_json" in current_block:
try:
current_block["input"] = json.loads(current_block["_partial_json"])
except json.JSONDecodeError:
current_block["input"] = {}
del current_block["_partial_json"]
response["content"].append(current_block)
current_block = None
elif event_type == "message_delta":
delta = data.get("delta", {})
if delta.get("stop_reason"):
response["stop_reason"] = delta["stop_reason"]
if data.get("usage"):
response["usage"].update(data["usage"])
return response if response.get("content") else None
def _response_to_sse(self, response: dict[str, Any], provider: str) -> list[bytes]:
"""Convert a response dict back to SSE format.
Args:
response: API response dict.
provider: Provider type for formatting.
Returns:
List of SSE event bytes.
"""
if provider != "anthropic":
return []
events: list[bytes] = []
# message_start
msg_start = {
"type": "message_start",
"message": {
"id": response.get("id", "msg_generated"),
"type": "message",
"role": response.get("role", "assistant"),
"model": response.get("model", "unknown"),
"content": [],
"stop_reason": None,
"usage": response.get("usage", {}),
},
}
events.append(f"event: message_start\ndata: {json.dumps(msg_start)}\n\n".encode())
# Content blocks
for idx, block in enumerate(response.get("content", [])):
# content_block_start
if block.get("type") == "text":
block_start = {
"type": "content_block_start",
"index": idx,
"content_block": {"type": "text", "text": ""},
}
elif block.get("type") == "tool_use":
block_start = {
"type": "content_block_start",
"index": idx,
"content_block": {
"type": "tool_use",
"id": block.get("id", f"toolu_{idx}"),
"name": block.get("name", ""),
"input": {},
},
}
else:
continue
events.append(
f"event: content_block_start\ndata: {json.dumps(block_start)}\n\n".encode()
)
# content_block_delta(s)
if block.get("type") == "text" and block.get("text"):
delta = {
"type": "content_block_delta",
"index": idx,
"delta": {"type": "text_delta", "text": block["text"]},
}
events.append(f"event: content_block_delta\ndata: {json.dumps(delta)}\n\n".encode())
elif block.get("type") == "tool_use" and block.get("input"):
delta = {
"type": "content_block_delta",
"index": idx,
"delta": {
"type": "input_json_delta",
"partial_json": json.dumps(block["input"]),
},
}
events.append(f"event: content_block_delta\ndata: {json.dumps(delta)}\n\n".encode())
# content_block_stop
block_stop = {"type": "content_block_stop", "index": idx}
events.append(f"event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n".encode())
# message_delta
msg_delta = {
"type": "message_delta",
"delta": {"stop_reason": response.get("stop_reason", "end_turn")},
"usage": {"output_tokens": response.get("usage", {}).get("output_tokens", 0)},
}
events.append(f"event: message_delta\ndata: {json.dumps(msg_delta)}\n\n".encode())
# message_stop
events.append(b'event: message_stop\ndata: {"type": "message_stop"}\n\n')
return events
def _record_ccr_feedback_from_response(
self, response: dict, provider: str, request_id: str
) -> None:
"""Extract headroom_retrieve tool calls from a response and record feedback.
This closes the TOIN feedback loop for streaming responses where
the proxy can't intercept and handle retrieval calls inline.
"""
from headroom.cache.compression_store import get_compression_store
content = response.get("content", [])
if not isinstance(content, list):
return
store = get_compression_store()
for block in content:
if not isinstance(block, dict):
continue
if block.get("type") != "tool_use":
continue
if block.get("name") != "headroom_retrieve":
continue
input_data = block.get("input", {})
hash_key = input_data.get("hash")
query = input_data.get("query")
if not hash_key:
continue
logger.info(
f"[{request_id}] CCR Feedback: Recording retrieval "
f"hash={hash_key[:8]}... query={query!r}"
)
# Call store.retrieve()/search() for the side effect of triggering
# the feedback chain: _log_retrieval -> process_pending_feedback
# -> toin.record_retrieval(). We discard the returned content.
try:
if query:
store.search(hash_key, query)
else:
store.retrieve(hash_key, query=None)
except Exception as e:
logger.debug(f"[{request_id}] CCR Feedback recording failed: {e}")
async def _stream_response(
self,
url: str,
headers: dict,
body: dict,
provider: str,
model: str,
request_id: str,
original_tokens: int,
optimized_tokens: int,
tokens_saved: int,
transforms_applied: list[str],
tags: dict[str, str],
optimization_latency: float,
memory_user_id: str | None = None,
pipeline_timing: dict[str, float] | None = None,
prefix_tracker: Any | None = None,
) -> StreamingResponse:
"""Stream response with metrics tracking and memory tool handling.
Parses SSE events to extract actual usage information from the API response
for accurate token counting and cost calculation.
When memory is enabled (memory_user_id provided), this method:
1. Buffers the response to detect memory tool calls
2. Executes memory tools if found
3. Makes continuation requests until no memory tools remain
4. Streams the final response to the client
"""
from fastapi.responses import StreamingResponse
from headroom.proxy.cost import _summarize_transforms
from headroom.proxy.helpers import MAX_SSE_BUFFER_SIZE
start_time = time.time()
# Mutable state for the generator to update
stream_state: dict[str, Any] = {
"input_tokens": None,
"output_tokens": None,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"total_bytes": 0,
"sse_buffer": "", # Buffer for incomplete SSE events
"ttfb_ms": None, # Time to first byte from upstream
}
# Track if we need to handle memory tools
memory_enabled = (
memory_user_id is not None
and self.memory_handler is not None
and provider == "anthropic"
)
# Open connection before generator to capture upstream response headers
# (needed to forward ratelimit headers to the client via StreamingResponse)
assert self.http_client is not None, "http_client must be initialized before streaming"
try:
_upstream_req = self.http_client.build_request("POST", url, json=body, headers=headers)
upstream_response = await self.http_client.send(_upstream_req, stream=True)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.PoolTimeout) as e:
error_msg = str(e)
logger.error(f"[{request_id}] Connection error to upstream API: {error_msg}")
async def _error_gen():
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": f"Failed to connect to upstream API: {error_msg}",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return StreamingResponse(_error_gen(), media_type="text/event-stream")
# Forward upstream ratelimit headers to the client
forwarded_headers = {
k: v for k, v in upstream_response.headers.items() if "ratelimit" in k.lower()
}
async def generate():
nonlocal body, memory_enabled # May need to modify for continuation requests
# For memory mode, we buffer the response to check for tool calls
buffered_chunks: list[bytes] = []
full_sse_data = ""
try:
async with contextlib.aclosing(upstream_response) as response:
async for chunk in response.aiter_bytes():
# Record TTFB on first chunk
if stream_state["ttfb_ms"] is None:
stream_state["ttfb_ms"] = (time.time() - start_time) * 1000
stream_state["total_bytes"] += len(chunk)
# Buffer SSE data to handle chunks split across calls
chunk_str = chunk.decode("utf-8", errors="ignore")
stream_state["sse_buffer"] += chunk_str
# Safety: prevent unbounded buffer growth
if len(stream_state["sse_buffer"]) > MAX_SSE_BUFFER_SIZE:
logger.error(
"SSE buffer exceeded maximum size (%d bytes), "
"truncating to prevent memory exhaustion",
MAX_SSE_BUFFER_SIZE,
)
stream_state["sse_buffer"] = stream_state["sse_buffer"][
-MAX_SSE_BUFFER_SIZE // 2 :
]
# Always stream immediately — buffering breaks
# real-time clients (LangGraph, LangChain, etc.)
yield chunk
if memory_enabled:
# Also buffer for post-stream memory processing
buffered_chunks.append(chunk)
full_sse_data += chunk_str
if len(full_sse_data) > MAX_SSE_BUFFER_SIZE:
logger.warning(
"Memory-mode SSE buffer exceeded maximum size, "
"disabling memory detection for this request"
)
memory_enabled = False
# Parse complete SSE events from buffer
usage = self._parse_sse_usage_from_buffer(stream_state, provider)
if usage:
if "input_tokens" in usage:
stream_state["input_tokens"] = usage["input_tokens"]
if "output_tokens" in usage:
stream_state["output_tokens"] = usage["output_tokens"]
if "cache_read_input_tokens" in usage:
stream_state["cache_read_input_tokens"] = usage[
"cache_read_input_tokens"
]
if "cache_creation_input_tokens" in usage:
stream_state["cache_creation_input_tokens"] = usage[
"cache_creation_input_tokens"
]
# Memory tool handling after stream completes
# Chunks were already yielded in real-time above, so we only
# do silent background processing here — no yielding.
if memory_enabled and full_sse_data:
# Check for Claude Code credential error
if "only authorized for use with Claude Code" in full_sse_data:
logger.warning(
f"[{request_id}] Memory: Claude Code subscription credentials "
"do not support custom tool injection. Set ANTHROPIC_API_KEY "
"environment variable or use --no-memory-tools flag."
)
return
# Parse SSE to get response JSON
parsed_response = self._parse_sse_to_response(full_sse_data, provider)
if parsed_response and self.memory_handler.has_memory_tool_calls(
parsed_response, provider
):
logger.info(
f"[{request_id}] Memory: Detected tool calls in streaming response"
)
# Execute memory tool calls silently — response already
# streamed so we cannot make a continuation request.
tool_results = await self.memory_handler.handle_memory_tool_calls(
parsed_response, memory_user_id, provider
)
if tool_results:
logger.info(
f"[{request_id}] Memory: Tool calls executed silently "
"(streaming mode — no continuation)"
)
# CCR Feedback: Record headroom_retrieve tool calls for TOIN learning.
# In streaming mode, the client handles actual retrieval, but we
# still need to record the event so TOIN learns which fields matter.
if self.config.ccr_inject_tool and full_sse_data:
ccr_parsed = (
parsed_response
if parsed_response
else self._parse_sse_to_response(full_sse_data, provider)
)
if ccr_parsed:
self._record_ccr_feedback_from_response(ccr_parsed, provider, request_id)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.PoolTimeout) as e:
logger.error(f"[{request_id}] Connection error to upstream API: {e}")
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": f"Failed to connect to upstream API: {e}",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.HTTPStatusError as e:
logger.error(f"[{request_id}] HTTP error from upstream API: {e}")
# Forward the upstream error response
yield e.response.content
except Exception as e:
logger.error(f"[{request_id}] Unexpected streaming error: {e}")
error_event = {
"type": "error",
"error": {"type": "api_error", "message": str(e)},
}
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"]
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:
prefix_tracker.update_from_response(
cache_read_tokens=cache_read_tokens,
cache_write_tokens=cache_write_tokens,
messages=body.get("messages", []),
)
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,
uncached_tokens=uncached_input_tokens,
)
await self.metrics.record_request(
provider=provider,
model=model,
input_tokens=optimized_tokens, # What we sent, not API's non-cached count
output_tokens=output_tokens,
tokens_saved=tokens_saved,
latency_ms=total_latency,
overhead_ms=optimization_latency,
ttfb_ms=stream_state["ttfb_ms"] or 0,
pipeline_timing=pipeline_timing,
cache_read_tokens=cache_read_tokens,
cache_write_tokens=cache_write_tokens,
uncached_input_tokens=uncached_input_tokens,
)
return StreamingResponse(
generate(),
media_type="text/event-stream",
headers=forwarded_headers,
)
async def _stream_response_bedrock(
self,
body: dict,
headers: dict,
provider: str,
model: str,
request_id: str,
original_tokens: int,
optimized_tokens: int,
tokens_saved: int,
transforms_applied: list[str],
tags: dict[str, str],
optimization_latency: float,
pipeline_timing: dict[str, float] | None = None,
) -> StreamingResponse:
"""Stream response from Bedrock backend with metrics tracking.
Translates Bedrock streaming events to Anthropic SSE format.
"""
from fastapi.responses import StreamingResponse
from headroom.proxy.cost import _summarize_transforms
from headroom.proxy.models import RequestLog
start_time = time.time()
# Mutable state for the generator
stream_state: dict[str, Any] = {
"input_tokens": 0,
"output_tokens": 0,
"ttfb_ms": None,
}
async def generate():
try:
assert self.anthropic_backend is not None
async for event in self.anthropic_backend.stream_message(body, headers):
# Record TTFB on first event
if stream_state["ttfb_ms"] is None:
stream_state["ttfb_ms"] = (time.time() - start_time) * 1000
# Format as SSE
if event.raw_sse:
yield event.raw_sse.encode()
else:
sse_line = f"event: {event.event_type}\ndata: {json.dumps(event.data)}\n\n"
yield sse_line.encode()
# Track usage from message_start event
if event.event_type == "message_start":
msg = event.data.get("message", {})
usage = msg.get("usage", {})
if "input_tokens" in usage:
stream_state["input_tokens"] = usage["input_tokens"]
# Track output tokens from message_delta
if event.event_type == "message_delta":
usage = event.data.get("usage", {})
if "output_tokens" in usage:
stream_state["output_tokens"] = usage["output_tokens"]
# Handle errors
if event.event_type == "error":
logger.error(f"[{request_id}] Bedrock stream error: {event.data}")
except Exception as e:
logger.error(f"[{request_id}] Bedrock streaming error: {e}")
error_event = {
"type": "error",
"error": {"type": "api_error", "message": str(e)},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
finally:
# Record metrics
total_latency = (time.time() - start_time) * 1000
output_tokens = stream_state["output_tokens"]
_backend_name = (
self.anthropic_backend.name if self.anthropic_backend else "anthropic"
)
await self.metrics.record_request(
provider=_backend_name,
model=model,
input_tokens=optimized_tokens,
output_tokens=output_tokens,
tokens_saved=tokens_saved,
latency_ms=total_latency,
cached=False,
overhead_ms=optimization_latency,
ttfb_ms=stream_state["ttfb_ms"] or 0,
pipeline_timing=pipeline_timing,
)
if self.cost_tracker:
self.cost_tracker.record_tokens(model, tokens_saved, optimized_tokens)
# Log request
if self.logger:
self.logger.log(
RequestLog(
request_id=request_id,
timestamp=datetime.now().isoformat(),
provider=_backend_name,
model=model,
input_tokens_original=original_tokens,
input_tokens_optimized=optimized_tokens,
output_tokens=output_tokens,
tokens_saved=tokens_saved,
savings_percent=(tokens_saved / original_tokens * 100)
if original_tokens > 0
else 0,
optimization_latency_ms=optimization_latency,
total_latency_ms=total_latency,
tags=tags,
cache_hit=False,
transforms_applied=transforms_applied,
request_messages=body.get("messages")
if self.config.log_full_messages
else None,
)
)
# Structured perf log line for `headroom perf` analysis
num_msgs = len(body.get("messages", []))
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=0 cache_write=0 cache_hit_pct=0 "
f"opt_ms={optimization_latency:.0f} "
f"transforms={_summarize_transforms(transforms_applied)}"
)
return StreamingResponse(
generate(),
media_type="text/event-stream",
)
async def _stream_openai_via_backend(
self,
body: dict,
headers: dict,
model: str,
request_id: str,
start_time: float,
original_tokens: int,
optimized_tokens: int,
tokens_saved: int,
transforms_applied: list[str],
tags: dict[str, str],
optimization_latency: float,
pipeline_timing: dict[str, float] | None = None,
) -> StreamingResponse:
"""Stream OpenAI chat completion response from backend.
Routes stream:true requests through the backend's stream_openai_message(),
yielding SSE events to the client.
"""
from fastapi.responses import StreamingResponse
assert self.anthropic_backend is not None
async def generate():
try:
async for sse_chunk in self.anthropic_backend.stream_openai_message(body, headers):
yield sse_chunk.encode() if isinstance(sse_chunk, str) else sse_chunk
except Exception as e:
logger.error(f"[{request_id}] Backend streaming error: {e}")
error_data = {
"error": {
"message": str(e),
"type": "api_error",
"code": "backend_error",
}
}
yield f"data: {json.dumps(error_data)}\n\n".encode()
yield b"data: [DONE]\n\n"
finally:
total_latency = (time.time() - start_time) * 1000
await self.metrics.record_request(
provider=self.anthropic_backend.name,
model=model,
input_tokens=optimized_tokens,
output_tokens=0, # Unknown in streaming
tokens_saved=tokens_saved,
latency_ms=total_latency,
cached=False,
overhead_ms=optimization_latency,
pipeline_timing=pipeline_timing,
)
if tokens_saved > 0:
logger.info(
f"[{request_id}] {model}: {original_tokens:,}{optimized_tokens:,} "
f"(saved {tokens_saved:,} tokens) via {self.anthropic_backend.name} [stream]"
)
return StreamingResponse(
generate(),
media_type="text/event-stream",
)

File diff suppressed because it is too large Load diff

View file

@ -272,6 +272,12 @@ module = [
disallow_untyped_defs = false
warn_return_any = false
# Handler mixins use self.* from HeadroomProxy via duck typing — mypy can't resolve these
[[tool.mypy.overrides]]
module = ["headroom.proxy.handlers.*"]
disallow_untyped_defs = false
ignore_errors = true
# Ignore third-party stubs with syntax errors
[[tool.mypy.overrides]]
module = ["mlx.*"]