Compare commits

...

5 Commits

Author SHA1 Message Date
Shroominic
6b9917c0bd max cost filtered models 2025-12-24 12:08:39 +01:00
Shroominic
b18714e092 ruff format 2025-12-22 19:31:37 +01:00
Shroominic
4181adaa23 rename responses to responses_api 2025-12-22 19:30:58 +01:00
9qeklajc
42546fd387 refactor 2025-12-15 21:48:15 +01:00
9qeklajc
cf28088afe clean up 2025-12-15 21:20:36 +01:00
11 changed files with 1305 additions and 1839 deletions

View File

@@ -133,8 +133,16 @@ async def calculate_cost( # todo: can be sync
output_tokens = response_data.get("usage", {}).get("completion_tokens", 0)
# added for response api
input_tokens = input_tokens if input_tokens != 0 else response_data.get("usage", {}).get("input_tokens", 0)
output_tokens = output_tokens if output_tokens != 0 else response_data.get("usage", {}).get("output_tokens", 0)
input_tokens = (
input_tokens
if input_tokens != 0
else response_data.get("usage", {}).get("input_tokens", 0)
)
output_tokens = (
output_tokens
if output_tokens != 0
else response_data.get("usage", {}).get("output_tokens", 0)
)
input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)

View File

@@ -3,14 +3,15 @@ import json
import random
import httpx
from fastapi import APIRouter, Depends
from fastapi import APIRouter, Depends, Request
from pydantic.v1 import BaseModel
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, create_session, get_session
from ..core.db import ApiKey, ModelRow, create_session, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from ..wallet import deserialize_token_from_string
from .price import sats_usd_price
logger = get_logger(__name__)
@@ -521,11 +522,89 @@ def _pricing_matches(
return True
async def _get_request_balance(request: Request, session: AsyncSession) -> int | None:
"""Get the balance from the request headers if authentication is provided."""
headers = request.headers
token: str | None = None
if x_cashu := headers.get("x-cashu"):
token = x_cashu
elif auth := headers.get("authorization"):
parts = auth.split(" ")
if len(parts) > 1:
token = parts[1]
if not token:
return None
# Handle API keys (sk-*)
if token.startswith("sk-"):
try:
# sk- keys use the part after "sk-" as the ID
key_id = token[3:]
key = await session.get(ApiKey, key_id)
if key:
return key.balance - key.reserved_balance
except Exception as e:
logger.warning(f"Error checking API key balance: {e}")
return None
# Handle Cashu tokens
try:
token_obj = deserialize_token_from_string(token)
amount_msat = (
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
)
return amount_msat
except Exception as e:
logger.debug(f"Failed to deserialize cashu token for balance check: {e}")
return None
@models_router.get("/v1/models")
@models_router.get("/models", include_in_schema=False)
async def models(session: AsyncSession = Depends(get_session)) -> dict:
async def models(
request: Request, session: AsyncSession = Depends(get_session)
) -> dict:
"""Get all available models from all providers with database overrides applied."""
from ..proxy import get_unique_models
items = get_unique_models()
# Optional: Filter by user balance if authenticated
user_balance = await _get_request_balance(request, session)
if user_balance is not None:
filtered_items = []
# Calculate tolerance factor once
tol_factor = 1.0
if not settings.fixed_pricing and settings.tolerance_percentage > 0:
tol_factor = 1.0 - (settings.tolerance_percentage / 100.0)
fixed_cost = settings.fixed_cost_per_request * 1000
min_cost = settings.min_request_msat
for model in items:
model_cost_msats = 0
if settings.fixed_pricing:
model_cost_msats = max(min_cost, fixed_cost)
elif model.sats_pricing:
# Use model specific pricing
max_cost = (
model.sats_pricing.max_cost
* 1000
* tol_factor
)
model_cost_msats = max(min_cost, int(max_cost))
else:
# Fallback if no pricing found
model_cost_msats = max(min_cost, fixed_cost)
if model_cost_msats <= user_balance:
filtered_items.append(model)
items = filtered_items
return {"data": items}

View File

@@ -146,7 +146,7 @@ async def proxy(
request_body_dict = parse_request_body_json(request_body, path)
if is_responses_api:
model_id = extract_model_from_responses_request(request_body_dict)
model_id = extract_model_from_responses_api_request(request_body_dict)
else:
model_id = request_body_dict.get("model", "unknown")
@@ -175,7 +175,7 @@ async def proxy(
if x_cashu := headers.get("x-cashu", None):
if is_responses_api:
return await upstream.handle_x_cashu_responses(
return await upstream.handle_x_cashu_responses_api(
request, x_cashu, path, max_cost_for_model, model_obj
)
else:
@@ -205,7 +205,7 @@ async def proxy(
headers = upstream.prepare_headers(dict(request.headers))
if is_responses_api:
response = await upstream.forward_responses_request(
response = await upstream.forward_responses_api_request(
request,
path,
headers,
@@ -328,7 +328,7 @@ async def get_bearer_token_key(
raise
def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> str:
def extract_model_from_responses_api_request(request_body_dict: dict[str, Any]) -> str:
if model := request_body_dict.get("model"):
return model
@@ -341,7 +341,7 @@ def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> s
logger.warning(
"No model found in Responses API request",
extra={"body_keys": list(request_body_dict.keys())}
extra={"body_keys": list(request_body_dict.keys())},
)
return "unknown"

File diff suppressed because it is too large Load Diff

View File

View File

@@ -0,0 +1,583 @@
"""X-Cashu payment request handlers for different API types."""
from __future__ import annotations
import json
from typing import TYPE_CHECKING, Callable
import httpx
from fastapi import BackgroundTasks, Request
from fastapi.responses import Response, StreamingResponse
from ...core import get_logger
from ...wallet import send_token
from ..payment.cashu_handler import (
CashuTokenError,
create_cashu_error_response,
validate_and_redeem_token,
)
from ..processing.http_client import HttpForwarder
from ..processing.response_handler import ChatCompletionProcessor, ResponsesApiProcessor
if TYPE_CHECKING:
from ...payment.models import Model
logger = get_logger(__name__)
class BaseCashuHandler:
"""Base handler for X-Cashu payment requests."""
def __init__(self, base_url: str):
self.base_url = base_url
self.http_forwarder = HttpForwarder(base_url)
async def send_refund(self, amount: int, unit: str, mint: str | None = None) -> str:
"""Create and send a refund token to the user."""
logger.debug(
"Creating refund token",
extra={"amount": amount, "unit": unit, "mint": mint},
)
max_retries = 3
last_exception = None
for attempt in range(max_retries):
try:
refund_token = await send_token(amount, unit=unit, mint_url=mint)
logger.info(
"Refund token created successfully",
extra={
"amount": amount,
"unit": unit,
"mint": mint,
"attempt": attempt + 1,
"token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
return refund_token
except Exception as e:
last_exception = e
if attempt < max_retries - 1:
logger.warning(
"Refund token creation failed, retrying",
extra={
"error": str(e),
"error_type": type(e).__name__,
"attempt": attempt + 1,
"max_retries": max_retries,
"amount": amount,
"unit": unit,
"mint": mint,
},
)
else:
logger.error(
"Failed to create refund token after all retries",
extra={
"error": str(e),
"error_type": type(e).__name__,
"attempt": attempt + 1,
"max_retries": max_retries,
"amount": amount,
"unit": unit,
"mint": mint,
},
)
raise Exception(
f"failed to create refund after {max_retries} attempts: {str(last_exception)}"
)
def _calculate_refund_amount(self, amount: int, unit: str, cost_msats: int) -> int:
"""Calculate refund amount based on unit and cost."""
if unit == "msat":
return amount - cost_msats
elif unit == "sat":
return amount - (cost_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
async def _handle_upstream_error(
self,
response: httpx.Response,
amount: int,
unit: str,
mint: str | None,
api_type: str = "API",
) -> Response:
"""Handle upstream service errors with refund."""
logger.warning(
f"Upstream {api_type} request failed, processing refund",
extra={
"status_code": response.status_code,
"amount": amount,
"unit": unit,
},
)
refund_token = await self.send_refund(amount - 60, unit, mint)
logger.info(
f"Refund processed for failed upstream {api_type} request",
extra={
"status_code": response.status_code,
"refund_amount": amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
error_response = Response(
content=json.dumps(
{
"error": {
"message": f"Error forwarding {api_type} request to upstream",
"type": "upstream_error",
"code": response.status_code,
"refund_token": refund_token,
}
}
),
status_code=response.status_code,
media_type="application/json",
)
error_response.headers["X-Cashu"] = refund_token
return error_response
class ChatCompletionCashuHandler(BaseCashuHandler):
"""Handler for X-Cashu paid chat completion requests."""
async def handle_request(
self,
request: Request,
x_cashu_token: str,
path: str,
headers: dict,
max_cost_for_model: int,
model_obj: Model,
prepare_request_body_func: Callable,
get_x_cashu_cost_func: Callable,
query_params_func: Callable,
) -> Response | StreamingResponse:
"""Handle chat completion request with X-Cashu payment."""
try:
amount, unit, mint = await validate_and_redeem_token(x_cashu_token, request)
# Forward request to upstream
request_body = await request.body()
response = await self.http_forwarder.forward_request(
request=request,
path=path,
headers=headers,
request_body=request_body,
query_params=query_params_func(path, request.query_params),
transform_body_func=prepare_request_body_func,
model_obj=model_obj,
)
# Handle upstream errors
if response.status_code != 200:
return await self._handle_upstream_error(
response, amount, unit, mint, "chat completion"
)
# Process chat completion response
if path.endswith("chat/completions"):
result = await self._handle_chat_completion_response(
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
result.background = background_tasks
return result
# Default streaming response for other endpoints
return self._create_default_streaming_response(response)
except CashuTokenError as e:
return create_cashu_error_response(e, request, x_cashu_token)
async def _handle_chat_completion_response(
self,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> Response | StreamingResponse:
"""Handle chat completion response processing."""
content = await response.aread()
content_str = content.decode("utf-8") if isinstance(content, bytes) else content
is_streaming = ChatCompletionProcessor.is_streaming_response(content_str)
if is_streaming:
return await self._handle_streaming_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
else:
return await self._handle_non_streaming_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
async def _handle_streaming_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> StreamingResponse:
"""Handle streaming chat completion response with refund calculation."""
response_headers = ChatCompletionProcessor.clean_response_headers(
dict(response.headers)
)
usage_data, model = ChatCompletionProcessor.extract_usage_from_streaming(
content_str
)
if usage_data and model:
try:
response_data = {"usage": usage_data, "model": model}
cost_data = await get_x_cashu_cost_func(
response_data, max_cost_for_model
)
if cost_data:
refund_amount = self._calculate_refund_amount(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(refund_amount, unit, mint)
response_headers["X-Cashu"] = refund_token
except Exception as e:
logger.error(
"Error calculating cost for streaming response",
extra={"error": str(e), "error_type": type(e).__name__},
)
lines = content_str.strip().split("\n")
return StreamingResponse(
ChatCompletionProcessor.create_streaming_generator(lines),
status_code=response.status_code,
headers=response_headers,
media_type="text/plain",
)
async def _handle_non_streaming_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> Response:
"""Handle non-streaming chat completion response with refund calculation."""
response_headers = ChatCompletionProcessor.clean_response_headers(
dict(response.headers)
)
try:
response_json = json.loads(content_str)
cost_data = await get_x_cashu_cost_func(response_json, max_cost_for_model)
if cost_data:
refund_amount = self._calculate_refund_amount(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(refund_amount, unit, mint)
response_headers["X-Cashu"] = refund_token
return Response(
content=content_str,
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError:
# Emergency refund on parse error
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response_headers["X-Cashu"] = refund_token
logger.warning(
"Emergency refund issued due to JSON parse error",
extra={
"original_amount": amount,
"refund_amount": emergency_refund,
"deduction": 60,
},
)
return Response(
content=content_str,
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
def _create_default_streaming_response(
self, response: httpx.Response
) -> StreamingResponse:
"""Create default streaming response for non-chat endpoints."""
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
return StreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=dict(response.headers),
background=background_tasks,
)
class ResponsesApiCashuHandler(BaseCashuHandler):
"""Handler for X-Cashu paid Responses API requests."""
async def handle_request(
self,
request: Request,
x_cashu_token: str,
path: str,
headers: dict,
max_cost_for_model: int,
model_obj: Model,
prepare_responses_api_request_body_func: Callable,
get_x_cashu_cost_func: Callable,
query_params_func: Callable,
) -> Response | StreamingResponse:
"""Handle Responses API request with X-Cashu payment."""
try:
amount, unit, mint = await validate_and_redeem_token(x_cashu_token, request)
# Forward request to upstream
request_body = await request.body()
response = await self.http_forwarder.forward_request(
request=request,
path=path,
headers=headers,
request_body=request_body,
query_params=query_params_func(path, request.query_params),
transform_body_func=prepare_responses_api_request_body_func,
model_obj=model_obj,
)
# Handle upstream errors
if response.status_code != 200:
return await self._handle_upstream_error(
response, amount, unit, mint, "Responses API"
)
# Process Responses API response
if path.startswith("responses"):
result = await self._handle_responses_api_completion(
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
result.background = background_tasks
return result
# Default streaming response for other endpoints
return self._create_default_streaming_response(response)
except CashuTokenError as e:
return create_cashu_error_response(e, request, x_cashu_token)
async def _handle_responses_api_completion(
self,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> Response | StreamingResponse:
"""Handle Responses API completion response processing."""
content = await response.aread()
content_str = content.decode("utf-8") if isinstance(content, bytes) else content
is_streaming = ResponsesApiProcessor.is_streaming_response(content_str)
if is_streaming:
return await self._handle_streaming_responses_api_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
else:
return await self._handle_non_streaming_responses_api_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
get_x_cashu_cost_func,
)
async def _handle_streaming_responses_api_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> StreamingResponse:
"""Handle streaming Responses API response with refund calculation."""
response_headers = ResponsesApiProcessor.clean_response_headers(
dict(response.headers)
)
# Extract usage data (ignoring reasoning tokens as per requirement)
usage_data, model = ResponsesApiProcessor.extract_usage_with_reasoning_tokens(
content_str
)
if usage_data and model:
try:
response_data = {"usage": usage_data, "model": model}
cost_data = await get_x_cashu_cost_func(
response_data, max_cost_for_model
)
if cost_data:
refund_amount = self._calculate_refund_amount(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(refund_amount, unit, mint)
response_headers["X-Cashu"] = refund_token
except Exception as e:
logger.error(
"Error calculating cost for streaming Responses API response",
extra={"error": str(e), "error_type": type(e).__name__},
)
lines = content_str.strip().split("\n")
return StreamingResponse(
ResponsesApiProcessor.create_streaming_generator(lines),
status_code=response.status_code,
headers=response_headers,
media_type="text/plain",
)
async def _handle_non_streaming_responses_api_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None,
get_x_cashu_cost_func: Callable,
) -> Response:
"""Handle non-streaming Responses API response with refund calculation."""
response_headers = ResponsesApiProcessor.clean_response_headers(
dict(response.headers)
)
try:
response_json = json.loads(content_str)
cost_data = await get_x_cashu_cost_func(response_json, max_cost_for_model)
if cost_data:
refund_amount = self._calculate_refund_amount(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(refund_amount, unit, mint)
response_headers["X-Cashu"] = refund_token
return Response(
content=content_str,
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError:
# Emergency refund on parse error
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response_headers["X-Cashu"] = refund_token
logger.warning(
"Emergency refund issued for Responses API due to JSON parse error",
extra={
"original_amount": amount,
"refund_amount": emergency_refund,
"deduction": 60,
},
)
return Response(
content=content_str,
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
def _create_default_streaming_response(
self, response: httpx.Response
) -> StreamingResponse:
"""Create default streaming response for non-responses endpoints."""
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
return StreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=dict(response.headers),
background=background_tasks,
)

View File

View File

@@ -0,0 +1,111 @@
"""X-Cashu payment handling utilities."""
from __future__ import annotations
from typing import TYPE_CHECKING
from fastapi import Request
from fastapi.responses import Response
from ...core import get_logger
from ...payment.helpers import create_error_response
from ...wallet import recieve_token
if TYPE_CHECKING:
pass
logger = get_logger(__name__)
class CashuTokenError(Exception):
"""Exception raised for X-Cashu token processing errors."""
def __init__(self, error_type: str, message: str, status_code: int = 400):
self.error_type = error_type
self.message = message
self.status_code = status_code
super().__init__(message)
async def validate_and_redeem_token(
x_cashu_token: str, request: Request
) -> tuple[int, str, str | None]:
"""Validate and redeem X-Cashu token.
Args:
x_cashu_token: The X-Cashu token to validate and redeem
request: Original FastAPI request (for error responses)
Returns:
Tuple of (amount, unit, mint_url)
Raises:
CashuTokenError: If token validation or redemption fails
"""
logger.info(
"Processing X-Cashu token redemption",
extra={
"token_preview": x_cashu_token[:20] + "..."
if len(x_cashu_token) > 20
else x_cashu_token,
},
)
try:
amount, unit, mint = await recieve_token(x_cashu_token)
logger.info(
"X-Cashu token redeemed successfully",
extra={"amount": amount, "unit": unit, "mint": mint},
)
return amount, unit, mint
except Exception as e:
error_message = str(e)
logger.error(
"X-Cashu token redemption failed",
extra={
"error": error_message,
"error_type": type(e).__name__,
},
)
# Determine specific error type
if "already spent" in error_message.lower():
raise CashuTokenError(
"token_already_spent",
"The provided CASHU token has already been spent",
400,
)
elif "invalid token" in error_message.lower():
raise CashuTokenError(
"invalid_token", "The provided CASHU token is invalid", 400
)
elif "mint error" in error_message.lower():
raise CashuTokenError(
"mint_error", f"CASHU mint error: {error_message}", 422
)
else:
raise CashuTokenError(
"cashu_error", f"CASHU token processing failed: {error_message}", 400
)
def create_cashu_error_response(
error: CashuTokenError, request: Request, x_cashu_token: str
) -> Response:
"""Create error response for X-Cashu token errors.
Args:
error: The CashuTokenError that occurred
request: Original FastAPI request
x_cashu_token: The token that caused the error
Returns:
Error response with appropriate status code and message
"""
return create_error_response(
error.error_type,
error.message,
error.status_code,
request=request,
token=x_cashu_token,
)

View File

View File

@@ -0,0 +1,200 @@
"""HTTP client utilities for upstream requests."""
from __future__ import annotations
import traceback
from typing import TYPE_CHECKING, Callable, Mapping
import httpx
from fastapi import BackgroundTasks, Request
from fastapi.responses import StreamingResponse
from ...core import get_logger
if TYPE_CHECKING:
from ...payment.models import Model
logger = get_logger(__name__)
class HttpForwarder:
"""Handles HTTP request forwarding to upstream services."""
def __init__(self, base_url: str):
self.base_url = base_url
def _prepare_path(self, path: str) -> str:
"""Prepare path by removing v1/ prefix if present."""
if path.startswith("v1/"):
path = path.replace("v1/", "")
return path
async def forward_request(
self,
request: Request,
path: str,
headers: dict,
request_body: bytes | None,
query_params: Mapping[str, str] | None,
transform_body_func: Callable | None = None,
model_obj: Model | None = None,
) -> httpx.Response:
"""Forward HTTP request to upstream service.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
request_body: Request body bytes, if any
query_params: Query parameters for the request
transform_body_func: Optional function to transform request body
model_obj: Model object for request body transformation
Returns:
Response from upstream service
Raises:
httpx.RequestError: If request fails
Exception: For other unexpected errors
"""
path = self._prepare_path(path)
url = f"{self.base_url}/{path}"
# Transform body if function provided
transformed_body = request_body
if transform_body_func and request_body and model_obj:
transformed_body = transform_body_func(request_body, model_obj)
logger.info(
"Forwarding request to upstream",
extra={
"url": url,
"method": request.method,
"path": path,
"has_request_body": request_body is not None,
},
)
client = httpx.AsyncClient(
transport=httpx.AsyncHTTPTransport(retries=1),
timeout=None,
)
try:
if transformed_body is not None:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=transformed_body,
params=query_params,
),
stream=True,
)
else:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=query_params,
),
stream=True,
)
logger.info(
"Received upstream response",
extra={
"status_code": response.status_code,
"path": path,
"content_type": response.headers.get("content-type", "unknown"),
},
)
return response
except httpx.RequestError as exc:
await client.aclose()
error_type = type(exc).__name__
error_details = str(exc)
logger.error(
"HTTP request error to upstream",
extra={
"error_type": error_type,
"error_details": error_details,
"method": request.method,
"url": url,
"path": path,
"query_params": dict(request.query_params),
},
)
if isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
elif isinstance(exc, httpx.TimeoutException):
error_message = "Upstream service request timed out"
elif isinstance(exc, httpx.NetworkError):
error_message = "Network error while connecting to upstream service"
else:
error_message = f"Error connecting to upstream service: {error_type}"
raise httpx.RequestError(error_message) from exc
except Exception as exc:
await client.aclose()
tb = traceback.format_exc()
logger.error(
"Unexpected error in upstream forwarding",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"method": request.method,
"url": url,
"path": path,
"query_params": dict(request.query_params),
"traceback": tb,
},
)
raise
class StreamingResponseWrapper:
"""Wrapper for creating streaming responses with proper cleanup."""
@staticmethod
def create_streaming_response(
response: httpx.Response,
client: httpx.AsyncClient,
) -> StreamingResponse:
"""Create a streaming response with background cleanup tasks.
Args:
response: httpx response to stream
client: httpx client to clean up
Returns:
StreamingResponse with background cleanup
"""
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
logger.debug(
"Creating streaming response",
extra={
"status_code": response.status_code,
"content_type": response.headers.get("content-type", "unknown"),
},
)
return StreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=dict(response.headers),
background=background_tasks,
)

View File

@@ -0,0 +1,170 @@
"""Response processing utilities for different API types."""
from __future__ import annotations
import json
import re
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING
from ...core import get_logger
if TYPE_CHECKING:
pass
logger = get_logger(__name__)
class ResponseProcessor:
"""Processes responses from upstream services."""
@staticmethod
def is_streaming_response(content_str: str) -> bool:
"""Determine if response content indicates streaming.
Args:
content_str: Response content as string
Returns:
True if response is streaming, False otherwise
"""
return content_str.startswith("data:") or "data:" in content_str
@staticmethod
def extract_usage_from_streaming(
content_str: str,
) -> tuple[dict | None, str | None]:
"""Extract usage data and model from streaming response content.
Args:
content_str: Streaming response content as string
Returns:
Tuple of (usage_data, model) or (None, None) if not found
"""
usage_data = None
model = None
lines = content_str.strip().split("\n")
for line in lines:
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if "usage" in data_json:
usage_data = data_json["usage"]
model = data_json.get("model")
elif "model" in data_json and not model:
model = data_json["model"]
except json.JSONDecodeError:
continue
return usage_data, model
@staticmethod
def clean_response_headers(headers: dict) -> dict:
"""Clean response headers by removing encoding-related headers.
Args:
headers: Original response headers
Returns:
Cleaned headers dict
"""
cleaned_headers = dict(headers)
headers_to_remove = ["transfer-encoding", "content-encoding", "content-length"]
for header in headers_to_remove:
cleaned_headers.pop(header, None)
return cleaned_headers
@staticmethod
async def create_streaming_generator(
lines: list[str],
) -> AsyncGenerator[bytes, None]:
"""Create async generator for streaming response lines.
Args:
lines: List of response lines to stream
Yields:
Encoded response lines
"""
for line in lines:
yield (line + "\n").encode("utf-8")
class ChatCompletionProcessor(ResponseProcessor):
"""Processor specifically for chat completion responses."""
@staticmethod
def extract_model_from_chunks(stored_chunks: list[bytes]) -> str | None:
"""Extract model name from stored response chunks.
Args:
stored_chunks: List of response chunks from streaming
Returns:
Model name if found, None otherwise
"""
last_model_seen = None
for i in range(len(stored_chunks) - 1, -1, -1):
chunk = stored_chunks[i]
if not chunk:
continue
try:
events = re.split(b"data: ", chunk)
for event_data in events:
if not event_data or event_data.strip() in (b"[DONE]", b""):
continue
try:
data = json.loads(event_data)
if isinstance(data, dict) and data.get("model"):
return str(data.get("model"))
except json.JSONDecodeError:
continue
except Exception as e:
logger.debug(
"Error processing chunk for model extraction",
extra={"error": str(e), "error_type": type(e).__name__},
)
return last_model_seen
class ResponsesApiProcessor(ResponseProcessor):
"""Processor specifically for Responses API responses."""
@staticmethod
def extract_usage_with_reasoning_tokens(
content_str: str,
) -> tuple[dict | None, str | None]:
"""Extract usage data including reasoning tokens from Responses API content.
Args:
content_str: Streaming response content as string
Returns:
Tuple of (usage_data, model) with reasoning tokens preserved
"""
usage_data = None
model = None
lines = content_str.strip().split("\n")
for line in lines:
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if "usage" in data_json:
usage_data = data_json["usage"]
model = data_json.get("model")
# Note: We now only track input_tokens and output_tokens
# reasoning_tokens are ignored per the requirement
break
elif "model" in data_json and not model:
model = data_json["model"]
except json.JSONDecodeError:
continue
return usage_data, model