mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
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:
parent
6a8ae297d6
commit
e8ab444f09
8 changed files with 5341 additions and 5084 deletions
20
headroom/proxy/handlers/__init__.py
Normal file
20
headroom/proxy/handlers/__init__.py
Normal 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",
|
||||
]
|
||||
1435
headroom/proxy/handlers/anthropic.py
Normal file
1435
headroom/proxy/handlers/anthropic.py
Normal file
File diff suppressed because it is too large
Load diff
995
headroom/proxy/handlers/batch.py
Normal file
995
headroom/proxy/handlers/batch.py
Normal 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)
|
||||
649
headroom/proxy/handlers/gemini.py
Normal file
649
headroom/proxy/handlers/gemini.py
Normal 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,
|
||||
}
|
||||
},
|
||||
)
|
||||
1320
headroom/proxy/handlers/openai.py
Normal file
1320
headroom/proxy/handlers/openai.py
Normal file
File diff suppressed because it is too large
Load diff
901
headroom/proxy/handlers/streaming.py
Normal file
901
headroom/proxy/handlers/streaming.py
Normal 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
|
|
@ -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.*"]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue