Files
routstr-core/routstr/upstream/base.py
T

5896 lines
231 KiB
Python

from __future__ import annotations
import asyncio
import json
import math
import traceback
import typing
import uuid
from collections.abc import AsyncGenerator, AsyncIterator, Iterator
from typing import Any, Mapping, Self, cast
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from pydantic.v1 import BaseModel
from ..auth import (
ReservationSnapshot,
adjust_payment_for_tokens,
get_reservation_snapshot,
release_reservation,
)
from ..core import get_logger
from ..core.db import (
ApiKey,
AsyncSession,
UpstreamProviderRow,
create_session,
)
from ..core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from ..core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
client_code_for_upstream_error,
client_status_for_upstream_error,
upstream_status_details,
)
from ..core.exceptions import UpstreamError
from ..core.redaction import redact_org_ids
from ..payment.cost_calculation import (
CostData,
CostDataError,
MaxCostData,
calculate_cost,
)
from ..payment.helpers import create_error_response
from ..payment.models import (
Model,
Pricing,
_calculate_usd_max_costs,
_update_model_sats_pricing,
backfill_cache_pricing,
list_models,
)
from ..payment.price import sats_usd_price
from ..wallet import (
SPENT_TOKEN_CODES,
classify_redemption_error,
recieve_token,
send_token,
token_mint_url,
)
from . import json_codec, messages_dispatch
from .cache_breakpoints import (
inject_anthropic_cache_breakpoints,
is_explicit_cache_model,
)
from .count_tokens import MissingUsageEstimator, count_tokens_locally
from .http_client import acquire_upstream_http_client, build_x_cashu_client
from .litellm_routing import detect_litellm_prefix
from .model_paths import public_provider_url
from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit
from .reasoning_effort import apply_reasoning_effort
from .sse_splitter import SSEEventSplitter
from .stream_ownership import (
ClosingStreamingResponse,
OwnedUpstreamStream,
PersistentStreamFinalizer,
ResponseHandoff,
aclose_if_needed,
attach_upstream_stream_owner,
close_upstream_exchange,
finalize_and_close_stream,
)
if typing.TYPE_CHECKING:
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
logger = get_logger(__name__)
CostMetadata = CostData | MaxCostData | dict[str, Any]
def _cost_field(
cost_data: CostMetadata, field: str, default: int | float = 0
) -> int | float:
if isinstance(cost_data, dict):
value = cost_data.get(field, default)
else:
value = getattr(cost_data, field, default)
return value if isinstance(value, (int, float)) else default
def _settled_cost_msats(cost_data: CostMetadata) -> int:
charged = _cost_field(cost_data, "charged_msats", -1)
if charged >= 0:
return int(charged)
return int(_cost_field(cost_data, "total_msats"))
def _published_cost(cost_data: CostMetadata) -> dict[str, Any]:
cost = dict(cost_data) if isinstance(cost_data, dict) else cost_data.dict()
computed_msats = int(_cost_field(cost_data, "total_msats"))
settled_msats = _settled_cost_msats(cost_data)
if computed_msats != settled_msats:
cost["computed_msats"] = computed_msats
cost["total_msats"] = settled_msats
cost["charged_msats"] = settled_msats
return cost
def _inject_cost_response_headers(
headers: dict[str, str], cost_data: CostMetadata
) -> None:
"""Inject per-request cost breakdown into response headers.
The SDK's ``extractUsageFromResponseHeaders`` reads these to populate
``inputMsats``, ``outputMsats``, ``totalMsats`` and ``satsCost`` in the
usage tracking entry — without them, x-cashu requests show 0.0 for all
sat cost fields.
"""
settled_msats = _settled_cost_msats(cost_data)
computed_msats = int(_cost_field(cost_data, "total_msats"))
headers["X-Routstr-Cost-Msats"] = str(settled_msats)
if computed_msats != settled_msats:
headers["X-Routstr-Computed-Cost-Msats"] = str(computed_msats)
headers["X-Routstr-Input-Cost-Msats"] = str(
int(_cost_field(cost_data, "input_msats"))
)
headers["X-Routstr-Output-Cost-Msats"] = str(
int(_cost_field(cost_data, "output_msats"))
)
total_usd = float(_cost_field(cost_data, "total_usd", 0.0))
if total_usd:
headers["X-Routstr-Cost-Usd"] = str(total_usd)
def _apply_estimated_usage(
response_json: dict[str, Any],
request_body: bytes | None,
model_obj: Model | None,
amount: int,
unit: str,
api: str,
) -> None:
"""Bill a buffered response from a local estimate when usage is missing."""
if response_json.get("usage"):
return
estimator = MissingUsageEstimator(request_body, model_obj)
estimator.observe(response_json)
estimated = estimator.estimated_usage(response_json.get("model"))
if not estimated:
return
logger.warning(
"No usage in non-streaming response, billing from local token estimate",
extra={
"api": api,
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated
def _parse_sse_events(content: str) -> list[tuple[list[str], str]]:
"""Split a buffered SSE body into ``(field_lines, data)`` pairs.
``data`` is the newline-joined payload the SSE spec reassembles from every
``data:`` line of one event, so multi-line JSON survives. Comment/keepalive
lines are dropped and events carrying no data at all are skipped; the
remaining ``event:``/``id:``/``retry:`` fields stay attached to their event
so Responses API framing is preserved on re-emission. A trailing event
without its blank-line terminator is still returned.
"""
events: list[tuple[list[str], str]] = []
normalized = content.replace("\r\n", "\n").replace("\r", "\n")
for raw_event in normalized.split("\n\n"):
field_lines: list[str] = []
data_lines: list[str] = []
for line in raw_event.split("\n"):
if line.startswith("data:"):
data_lines.append(line[len("data:") :].lstrip(" "))
elif line and not line.startswith(":"):
field_lines.append(line)
if not data_lines:
continue
events.append((field_lines, "\n".join(data_lines)))
return events
def _responses_usage_payload(data_json: dict) -> dict:
"""Return the object carrying a Responses API event's model and usage.
Canonical events nest them under ``response`` (``response.completed`` /
``response.incomplete``); legacy and compat shapes keep them at top level.
"""
nested = data_json.get("response")
return nested if isinstance(nested, dict) else data_json
def _reported_provider(payload: dict) -> str | None:
"""Provider named by an upstream payload, if any.
Checked at top level first, then inside the Anthropic ``message`` and
Responses ``response`` envelopes, which is where those dialects nest it.
"""
for obj in (payload, payload.get("message"), payload.get("response")):
if isinstance(obj, dict):
value = obj.get("provider")
if isinstance(value, str) and value.strip():
return value.strip()
return None
def _render_sse_event(field_lines: list[str], data: str) -> str:
"""Re-frame one parsed event, re-prefixing every line of a multi-line data."""
body = "".join(f"{line}\n" for line in field_lines)
body += "".join(f"data: {line}\n" for line in data.split("\n"))
return body + "\n"
def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> None:
"""Inject cost breakdown into the response body's ``usage.cost`` object.
The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be
an object with ``total_msats``/``input_msats``/``output_msats`` (not a
plain USD number). When the upstream returns ``cost`` as a number, the
SDK cannot extract the msats breakdown from the body alone.
"""
usage = response_json.get("usage")
if not isinstance(usage, dict):
return
# Direct assignment (not setdefault) so routstr's authoritative cost
# data always overwrites any upstream-provided cost values. Using
# setdefault would silently keep stale upstream values and drop our
# calculated msats breakdown.
computed_msats = int(_cost_field(cost_data, "total_msats"))
settled_msats = _settled_cost_msats(cost_data)
cost_obj: dict[str, int | float] = {
"base_msats": int(_cost_field(cost_data, "base_msats")),
"input_msats": int(_cost_field(cost_data, "input_msats")),
"output_msats": int(_cost_field(cost_data, "output_msats")),
"total_msats": settled_msats,
"charged_msats": settled_msats,
"cache_read_input_tokens": int(
_cost_field(cost_data, "cache_read_input_tokens")
),
"cache_creation_input_tokens": int(
_cost_field(cost_data, "cache_creation_input_tokens")
),
"cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")),
"cache_creation_msats": int(_cost_field(cost_data, "cache_creation_msats")),
}
if computed_msats != settled_msats:
cost_obj["computed_msats"] = computed_msats
total_usd = float(_cost_field(cost_data, "total_usd", 0.0))
if total_usd:
cost_obj["total_usd"] = total_usd
usage["cost"] = cost_obj
usage["cost_sats"] = settled_msats // 1000
def _is_json_content_type(content_type: str | None) -> bool:
"""Return True when the upstream response should be parsed as JSON."""
if not content_type:
return False
main = content_type.split(";", 1)[0].strip().lower()
if main in ("application/json", "text/json"):
return True
return main.startswith("application/") and main.endswith("+json")
def _is_sse_body(content_type: str | None, content_str: str) -> bool:
if content_type:
main = content_type.split(";", 1)[0].strip().lower()
if main == "text/event-stream":
return True
if _is_json_content_type(content_type):
return False
for line in content_str.lstrip("\ufeff").splitlines():
stripped = line.strip()
if stripped:
return stripped.startswith(("data:", "event:", "id:", "retry:", ":"))
return False
def _openai_completion_path(path: str) -> str | None:
canonical = "/" + path.rstrip("/")
if canonical.endswith("/chat/completions"):
return "chat/completions"
return "completions" if canonical.endswith("/completions") else None
def _x_cashu_path_has_settlement_handler(path: str) -> bool:
canonical = path.rstrip("/")
return _openai_completion_path(canonical) is not None or canonical.endswith(
("embeddings", "messages", "messages/count_tokens", "systemone", "decisions")
)
class TopupData(BaseModel):
"""Universal top-up data schema for Lightning Network invoices."""
invoice_id: str
payment_request: str
amount: int
currency: str
expires_at: int | None = None
checkout_url: str | None = None
class BaseUpstreamProvider:
"""Provider for forwarding requests to an upstream AI service API."""
provider_type: str = "base"
default_base_url: str | None = None
platform_url: str | None = None
supports_anthropic_messages: bool = False
supports_decisions: bool = False
# When None, the prefix is detected from `base_url` at dispatch time
# (see `get_litellm_provider_prefix`). Subclasses set this to lock the
# provider regardless of URL.
litellm_provider_prefix: str | None = None
base_url: str
api_key: str
provider_fee: float = 1.05
# Primary key of the ``upstream_providers`` row this instance was built
# from. Set by ``from_db_row`` so a live provider can re-find its own row by
# stable identity instead of its rotatable ``api_key``. ``None`` for
# instances not sourced from a row.
db_id: int | None = None
_models_cache: list[Model] = []
_models_by_id: dict[str, Model] = {}
def __init__(self, base_url: str, api_key: str, provider_fee: float = 1.01):
"""Initialize the upstream provider.
Args:
base_url: Base URL of the upstream API endpoint
api_key: API key for authenticating with the upstream service
provider_fee: Provider fee multiplier (default 1.01 for 1% fee)
"""
self.base_url = base_url
self.api_key = api_key
self.provider_fee = provider_fee
self.db_id = None
self._models_cache = []
self._models_by_id = {}
def get_litellm_provider_prefix(self) -> str:
"""Resolve the litellm provider prefix for this provider instance.
1. If the subclass pinned `litellm_provider_prefix`, use it.
2. Otherwise infer from `base_url` (e.g. ``api.fireworks.ai`` →
``fireworks_ai/``) so custom/generic rows reach the correct
litellm backend instead of falling back to ``openai/``.
3. Default ``openai/`` for unknown OpenAI-compatible servers.
"""
if self.__class__.litellm_provider_prefix:
return self.__class__.litellm_provider_prefix
return detect_litellm_prefix(self.base_url)
@classmethod
def from_db_row(cls, provider_row: "UpstreamProviderRow") -> "Self | None":
"""Instantiate a provider from a database row, carrying its identity.
Construction itself is delegated to the ``_build_from_row`` hook (which
subclasses override to match their constructor); this wrapper stamps the
row's primary key onto the instance as ``db_id`` so the provider can
later re-find its own row by identity rather than by its ``api_key``.
Args:
provider_row: Database row containing provider configuration
Returns:
Instantiated provider or None if instantiation fails
"""
provider = cls._build_from_row(provider_row)
if provider is not None:
provider.db_id = provider_row.id
return provider
@classmethod
def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "Self | None":
"""Construct the provider instance from a row (no identity stamping).
Overridden by subclasses whose constructors differ from the base
``(base_url, api_key, provider_fee)`` shape. Callers should use
``from_db_row`` instead, which also attaches ``db_id``.
"""
return cls(
base_url=provider_row.base_url,
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
)
@classmethod
def get_provider_metadata(cls) -> dict[str, object]:
"""Get metadata about this provider type for API responses.
Returns:
Dict with provider type metadata including id, name, default_base_url, fixed_base_url, platform_url, can_create_account, can_topup, can_show_balance
"""
return {
"id": cls.provider_type,
"name": cls.provider_type.title(),
"default_base_url": cls.default_base_url or "",
"fixed_base_url": bool(cls.default_base_url),
"platform_url": cls.platform_url,
"can_create_account": False,
"can_topup": False,
"can_show_balance": False,
}
@staticmethod
def _fold_cache_into_input_tokens(usage: object) -> None:
"""Fold cache token counts into ``input_tokens`` / ``prompt_tokens``.
Cost calculation has already used the per-bucket counts to bill the
request correctly; what the client sees in the visible token total
should be a single rolled-up prompt count *including* the cache
portion. The standalone ``cache_read_input_tokens`` /
``cache_creation_input_tokens`` fields are left in place for clients
that want the breakdown.
Which field may be folded mirrors ``normalize_usage`` exactly:
Anthropic-native ``input_tokens`` *excludes* the cached portion and
needs the roll-up, while a ``prompt_tokens`` grand total (OpenAI
family, DeepSeek, OpenRouter, litellm) *already includes* it —
folding there double-counts the cache in the visible prompt total
(Venice showed 27997 prompt tokens for a 14075-token prompt after a
13922-token cache read).
For Anthropic-shaped responses (``input_tokens`` present), the cache
fields are forced to ``0`` when the upstream omitted them, so the
client always sees a consistent shape.
"""
if not isinstance(usage, dict):
return
# Normalise missing cache fields to 0 on Anthropic-shaped usage so
# downstream consumers can rely on them being present.
if "input_tokens" in usage:
usage.setdefault("cache_read_input_tokens", 0)
usage.setdefault("cache_creation_input_tokens", 0)
try:
cache_read = int(usage.get("cache_read_input_tokens") or 0)
cache_creation = int(usage.get("cache_creation_input_tokens") or 0)
except (TypeError, ValueError):
return
extra = cache_read + cache_creation
if extra <= 0:
return
if "input_tokens" in usage:
try:
usage["input_tokens"] = int(usage.get("input_tokens") or 0) + extra
except (TypeError, ValueError):
pass
# ``prompt_tokens`` is deliberately left untouched: in every dialect
# that reports it, it is an inclusive grand total that already
# contains the cached portion — the same assumption
# ``normalize_usage`` subtracts against when billing.
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the routstr ``provider`` field onto an upstream response payload.
Format is ``"<provider_type>:<upstream_provider>"`` when the upstream
already reported its own provider (e.g. OpenRouter returns
``"provider": "Fireworks"``), otherwise just ``"<provider_type>"``
for direct upstreams.
Idempotent: re-stamping an already-stamped payload must not nest the
prefix repeatedly (e.g. never ``"anthropic:anthropic"``). This matters
because streaming paths can apply the field more than once per chunk.
Also stamps ``provider_url`` with the upstream base URL that served
the request.
"""
if not isinstance(response_json, dict):
return
response_json["provider_url"] = public_provider_url(self.base_url)
provider_type = (self.provider_type or "").strip()
existing_str = _reported_provider(response_json) or ""
if not existing_str:
response_json["provider"] = provider_type
return
# Already stamped by a previous pass — leave it untouched.
if existing_str == provider_type or existing_str.startswith(
f"{provider_type}:"
):
response_json["provider"] = existing_str
return
response_json["provider"] = f"{provider_type}:{existing_str}"
def _stamp_streamed_provider(
self, payload: dict, carried: str | None
) -> str | None:
"""Stamp a streamed payload, falling back to a provider an earlier event
reported. Returns the provider to carry forward to later payloads."""
reported = _reported_provider(payload)
if reported is None and carried is not None:
payload["provider"] = carried
self._apply_provider_field(payload)
return reported or carried
def _log_full_refund(
self,
*,
route: str,
model: str | None,
content_str: str,
amount: int,
unit: str,
) -> None:
"""Record a settlement that serves content but charges nothing.
The client keeps both the response and the whole prepayment, so the
model, the serving upstream and a redacted body preview are logged to
keep the unbilled request auditable.
"""
logger.warning(
"Zero-cost settlement, refunding the full prepayment",
extra={
"route": route,
"model": model or "unknown",
"provider_type": self.provider_type,
"upstream_base_url": self.base_url,
"refund_amount": amount,
"unit": unit,
"response_body_preview": redact_org_ids(content_str.strip()[:500]),
},
)
def inject_cost_metadata(
self,
response_json: dict,
cost_data: CostData | MaxCostData | dict,
key: ApiKey,
) -> None:
"""Unifies the injection of cost and usage metadata across all completion types."""
self._apply_provider_field(response_json)
cost_dict = _published_cost(cost_data)
sats_cost = cost_dict["total_msats"] // 1000
# Inject the shared SDK cost contract into every usage shape.
if isinstance(response_json.get("usage"), dict):
_inject_cost_into_usage(response_json, cost_data)
response_json["usage"]["remaining_balance_msats"] = key.balance
self._fold_cache_into_input_tokens(response_json["usage"])
message = response_json.get("message")
if isinstance(message, dict) and isinstance(message.get("usage"), dict):
_inject_cost_into_usage(message, cost_data)
message["usage"]["remaining_balance_msats"] = key.balance
self._fold_cache_into_input_tokens(message["usage"])
nested_response = response_json.get("response")
if isinstance(nested_response, dict) and isinstance(
nested_response.get("usage"), dict
):
_inject_cost_into_usage(nested_response, cost_data)
nested_response["usage"]["remaining_balance_msats"] = key.balance
self._fold_cache_into_input_tokens(nested_response["usage"])
# Unified Routstr metadata
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {
"cost": cost_dict,
"sats_cost": sats_cost,
"remaining_balance_msats": key.balance,
}
# Legacy/Compatibility fields
response_json["cost"] = cost_dict.copy()
response_json["cost"]["sats_cost"] = sats_cost
response_json["cost"]["remaining_balance_msats"] = key.balance
def prepare_headers(self, request_headers: dict) -> dict:
"""Prepare headers for upstream request by removing proxy-specific headers and adding authentication.
Args:
request_headers: Original request headers from the client
Returns:
Headers dict ready for upstream forwarding with authentication added
"""
logger.debug(
"Preparing upstream headers",
extra={
"original_headers_count": len(request_headers),
"has_upstream_api_key": bool(self.api_key),
},
)
headers = dict(request_headers)
removed_headers = []
for header in [
"host",
"content-length",
"refund-lnurl",
"key-expiry-time",
"x-cashu",
"x-routstr-model-path",
]:
if headers.pop(header, None) is not None:
removed_headers.append(header)
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
if headers.pop("authorization", None) is not None:
removed_headers.append("authorization (replaced with upstream key)")
else:
for auth_header in ["Authorization", "authorization"]:
if headers.pop(auth_header, None) is not None:
removed_headers.append(auth_header)
for header in ["authorization", "accept-encoding"]:
if headers.pop(header, None) is not None:
removed_headers.append(f"{header} (replaced with routstr-safe version)")
# Explicitly define the list of supported compression encodings
headers["accept-encoding"] = "gzip, deflate, br, identity"
logger.debug(
"Headers prepared for upstream",
extra={
"final_headers_count": len(headers),
"removed_headers": removed_headers,
"added_upstream_auth": bool(self.api_key),
},
)
return headers
def prepare_params(
self, path: str, query_params: Mapping[str, str] | None
) -> Mapping[str, str]:
"""Prepare query parameters for upstream request.
Base implementation passes through query params unchanged. Override in subclasses for provider-specific params.
Args:
path: Request path
query_params: Original query parameters from the client
Returns:
Query parameters dict ready for upstream forwarding
"""
return query_params or {}
def transform_model_name(self, model_id: str) -> str:
"""Transform model ID for this provider's API format.
Base implementation returns model_id unchanged. Override in subclasses for provider-specific transformations.
Args:
model_id: Model identifier (may include provider prefix)
Returns:
Transformed model ID for this provider
"""
return model_id
def normalize_request_path(self, path: str, model_obj: Model | None = None) -> str:
"""Normalize request path before forwarding to upstream."""
if path.startswith("v1/"):
return path.replace("v1/", "", 1)
return path
def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str:
"""Get upstream base URL used when building forwarding URL."""
return self.base_url.rstrip("/")
def build_request_url(self, path: str, model_obj: Model | None = None) -> str:
"""Build full upstream URL from normalized path."""
clean_path = path.lstrip("/")
return f"{self.get_request_base_url(path, model_obj)}/{clean_path}"
def prepare_responses_request_body(
self, body: bytes | None, model_obj: Model
) -> bytes | None:
"""Transform request body for Responses API specific requirements.
Handles Responses API specific transformations while maintaining model name transforms.
Args:
body: Original request body bytes
model_obj: Model object containing the original model information
Returns:
Transformed request body bytes
"""
if not body:
return body
try:
data = json.loads(body)
if isinstance(data, dict):
# Handle model transformation in various locations
if "model" in data:
original_model = model_obj.id
transformed_model = self.transform_model_name(original_model)
data["model"] = transformed_model
logger.debug(
"Transformed model name in Responses API request",
extra={
"original": original_model,
"transformed": transformed_model,
"provider": self.provider_type or self.base_url,
},
)
# Handle model in input field (alternative format)
if (
"input" in data
and isinstance(data["input"], dict)
and "model" in data["input"]
):
original_model = model_obj.id
transformed_model = self.transform_model_name(original_model)
data["input"]["model"] = transformed_model
apply_reasoning_effort(data, model_obj)
return json.dumps(data).encode()
except Exception as e:
logger.debug(
"Could not transform Responses API request body",
extra={
"error": str(e),
"provider": self.provider_type or self.base_url,
},
)
return body
def _upstream_accepts_cache_control(self) -> bool:
"""True when this upstream accepts explicit ``cache_control`` markers.
Only OpenRouter (documents Anthropic + Alibaba explicit caching) and the
native Anthropic API accept the markers. Stamping them toward an
automatic-cache or non-supporting upstream risks a 400, so injection is
confined to these. Base URL is also checked so an OpenRouter endpoint
configured through the generic provider is still recognised.
"""
if self.provider_type in ("openrouter", "anthropic"):
return True
return "openrouter.ai" in (self.base_url or "")
def prepare_request_body(
self,
body: bytes | None,
model_obj: Model,
include_stream_usage: bool = False,
) -> bytes | None:
"""Transform request body for provider-specific requirements.
Automatically transforms model names and opts streaming OpenAI
completion endpoints into emitting per-chunk ``usage`` so cost
tracking can read real token counts.
Args:
body: Original request body bytes
include_stream_usage: Opt a streaming completion into usage chunks
Returns:
Transformed request body bytes
"""
if not body:
return body
try:
data = json.loads(body)
except Exception as e:
logger.debug(
"Could not parse request body for transformation",
extra={
"error": str(e),
"provider": self.provider_type or self.base_url,
},
)
return body
if not isinstance(data, dict):
return body
changed = False
if "model" in data:
original_model = model_obj.id
transformed_model = self.transform_model_name(original_model)
if data["model"] != transformed_model:
data["model"] = transformed_model
logger.debug(
"Transformed model name in request",
extra={
"original": original_model,
"transformed": transformed_model,
"provider": self.provider_type or self.base_url,
},
)
changed = True
# OpenAI-compatible streaming responses omit ``usage`` unless the
# request sets ``stream_options.include_usage = true``. Without it
# we can't reconcile token counts at end of stream and must use
# the local request/response estimator.
if data.get("stream") is True and include_stream_usage:
existing = data.get("stream_options")
merged = dict(existing) if isinstance(existing, dict) else {}
if merged.get("include_usage") is not True:
merged["include_usage"] = True
data["stream_options"] = merged
changed = True
# Explicit-cache models (Anthropic Claude, Alibaba Qwen / deepseek-v3.2)
# cache nothing without ``cache_control`` markers in the body. Clients
# that don't recognise a routstr URL as one of these never send them, so
# caching silently never engages over routstr even though it works
# against OpenRouter directly. Stamp the standard breakpoints so caching
# works by default, deferring to any client-set markers. Gated to
# upstreams that accept the markers (OpenRouter / Anthropic) so they
# never leak to an automatic-cache provider that would reject them.
if (
"messages" in data
and isinstance(data.get("messages"), list)
and self._upstream_accepts_cache_control()
and is_explicit_cache_model(
model_obj.id,
model_obj.forwarded_model_id,
model_obj.canonical_slug,
)
):
if inject_anthropic_cache_breakpoints(data):
changed = True
if apply_reasoning_effort(data, model_obj):
changed = True
if changed:
return json.dumps(data).encode()
return body
def _extract_upstream_error_message(
self, body_bytes: bytes
) -> tuple[str, str | None]:
"""Extract error message and code from upstream error response body.
Args:
body_bytes: Raw response body bytes from upstream
Returns:
Tuple of (error_message, error_code), where error_code may be None
"""
message: str = "Upstream request failed"
upstream_code: str | None = None
if not body_bytes:
return message, upstream_code
try:
data = json.loads(body_bytes)
if isinstance(data, dict):
err = data.get("error")
if isinstance(err, dict):
raw_msg = (
err.get("message") or err.get("detail") or err.get("error")
)
if isinstance(raw_msg, (str, int, float)):
message = str(raw_msg)
upstream_code_raw = err.get("code") or err.get("type")
if isinstance(upstream_code_raw, (str, int, float)):
upstream_code = str(upstream_code_raw)
elif "message" in data and isinstance(
data["message"], (str, int, float)
):
message = str(data["message"]) # type: ignore[arg-type]
elif "detail" in data and isinstance(data["detail"], (str, int, float)):
message = str(data["detail"]) # type: ignore[arg-type]
except Exception:
preview = body_bytes.decode("utf-8", errors="ignore").strip()
if preview:
message = preview[:500]
return redact_org_ids(message), upstream_code
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
) -> None:
"""Hook called when the proxy redirects to another provider due to an error.
Subclasses can implement this to perform actions like disabling the provider
if it's out of balance.
Args:
status_code: The HTTP status code returned by the upstream
error_message: The error message extracted from the upstream response
"""
pass
async def forward_upstream_error_response(
self,
request: Request,
path: str,
upstream_response: httpx.Response,
model_id: str | None = None,
) -> Response:
"""Log upstream errors and forward the response in a JSON envelope."""
status_code = upstream_response.status_code
headers = dict(upstream_response.headers)
content_type = headers.get("content-type") or headers.get("Content-Type", "")
upstream_request_id = (
headers.get("request-id")
or headers.get("Request-Id")
or headers.get("x-request-id")
or headers.get("X-Request-Id")
or headers.get("anthropic-request-id")
or headers.get("openai-request-id")
)
body_read_error = None
try:
body_bytes = await upstream_response.aread()
except Exception as exc:
body_bytes = b""
body_read_error = f"{type(exc).__name__}: {exc}"
# ``message`` is already redacted by ``_extract_upstream_error_message``;
# the raw body preview is redacted here before it reaches logs or the
# forwarded envelope so provider account identifiers never leak.
message, upstream_code = self._extract_upstream_error_message(body_bytes)
body_preview = redact_org_ids(
body_bytes.decode("utf-8", errors="ignore").strip()[:500]
)
is_json_body = _is_json_content_type(content_type)
# Classify upstream rate-limit failures into a stable, structured error.
rate_limit = classify_rate_limit(status_code, message, headers)
error_code: str | int = upstream_code or status_code
error_details: dict[str, object] | None = None
if rate_limit is not None:
error_code = UPSTREAM_RATE_LIMIT
error_details = rate_limit.as_details()
client_status = client_status_for_upstream_error(status_code, error_code)
client_code = client_code_for_upstream_error(status_code, error_code)
headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM
logger.warning(
"Upstream %s returned %s for model=%s path=%s: %s",
self.provider_type,
status_code,
model_id or "unknown",
path,
(message or body_preview or "<empty>")[:300],
extra={
"path": path,
"provider": self.provider_type,
"model": model_id or "unknown",
"upstream_status": status_code,
"upstream_code": upstream_code,
"error_code": error_code,
"upstream_content_type": content_type,
"upstream_request_id": upstream_request_id,
"message_preview": message[:200],
"body_preview": body_preview,
"body_read_error": body_read_error,
"method": request.method,
"json_normalized": not is_json_body,
},
)
for header_name in (
"content-length",
"Content-Length",
"transfer-encoding",
"Transfer-Encoding",
"content-encoding",
"Content-Encoding",
"connection",
"Connection",
"keep-alive",
"Keep-Alive",
"proxy-authenticate",
"Proxy-Authenticate",
"proxy-authorization",
"Proxy-Authorization",
"te",
"TE",
"trailer",
"Trailer",
"upgrade",
"Upgrade",
):
headers.pop(header_name, None)
# Propagate a usable retry hint to the caller when the upstream supplied
# one but did not echo a ``Retry-After`` header. RFC 7231 delta-seconds
# is an integer, so round sub-second hints up to a usable ``1``.
if (
rate_limit is not None
and rate_limit.retry_after_seconds is not None
and "retry-after" not in {k.lower() for k in headers}
):
headers["Retry-After"] = str(
max(1, math.ceil(rate_limit.retry_after_seconds))
)
if is_json_body:
if not content_type:
headers.pop("content-type", None)
headers.pop("Content-Type", None)
media_type = content_type or None
# Re-serialise the body with organization IDs stripped. The narrow
# ``org-*`` regex preserves the surrounding JSON structure.
redacted_text = redact_org_ids(body_bytes.decode("utf-8", errors="ignore"))
redacted_body = redacted_text.encode()
# Surface the stable classification on the forwarded body so callers
# can switch on ``error.code`` without parsing the provider-specific
# message. Fall back to the redacted bytes if the body is not a JSON
# object with an ``error`` mapping.
if rate_limit is not None or client_status != status_code:
try:
parsed = json.loads(redacted_text)
err = parsed.get("error") if isinstance(parsed, dict) else None
if isinstance(err, dict):
if rate_limit is not None:
err["code"] = UPSTREAM_RATE_LIMIT
err["details"] = error_details
if client_status != status_code:
err["code"] = client_code
err["upstream_status"] = status_code
redacted_body = json.dumps(parsed).encode()
elif (
client_status != status_code
and isinstance(parsed, dict)
and "error" not in parsed
):
# JSON body without an ``error`` mapping (e.g. FastAPI's
# ``{"detail": ...}``). Add one so a rewritten status is
# never served without its classification.
parsed["error"] = {
"message": message or "Upstream returned an error response",
"type": "upstream_error",
"code": client_code,
"upstream_status": status_code,
}
redacted_body = json.dumps(parsed).encode()
except (ValueError, AttributeError):
pass
return Response(
content=redacted_body,
status_code=client_status,
headers=headers,
media_type=media_type,
)
# Non-JSON upstream error (HTML, plain text, empty, ...). Wrap it in
# the standard JSON envelope so callers don't need a second parser.
for header_name in ("content-type", "Content-Type"):
headers.pop(header_name, None)
error_obj: dict[str, object] = {
"message": message or "Upstream returned a non-JSON error response",
"type": "upstream_error",
"code": client_code,
"upstream_status": status_code,
"upstream_content_type": content_type or None,
"upstream_body_preview": body_preview or None,
}
if error_details is not None:
error_obj["details"] = error_details
envelope = {
"error": error_obj,
"request_id": getattr(request.state, "request_id", None),
}
return Response(
content=json.dumps(envelope).encode(),
status_code=client_status,
headers=headers,
media_type="application/json",
)
async def _release_failed_streaming_reservation(
self,
key: ApiKey,
session: AsyncSession,
reservation_snapshot: ReservationSnapshot | None,
) -> bool:
"""Attempt exact release and suppress unsafe settlement retries."""
try:
await session.rollback()
snapshot = reservation_snapshot
if snapshot is None:
snapshot = await get_reservation_snapshot(key, session)
released = await release_reservation(
snapshot,
session,
snapshot.reserved_msats,
)
if not released:
logger.critical(
"Billing reservation could not be released",
extra={
"key_hash": key.hashed_key[:8] + "...",
"reserved_balance": snapshot.reserved_msats,
},
)
# A failed release remains recoverable by the stale-reservation
# sweep. Retrying settlement here could charge after an ambiguous
# database failure or replace the original stream exception.
return True
except asyncio.CancelledError:
# Preserve the exception that triggered billing cleanup. The stream
# propagates it immediately after this helper returns, and stale
# reservation cleanup can recover an interrupted release.
logger.critical(
"Billing reservation release was cancelled",
extra={"key_hash": key.hashed_key[:8] + "..."},
exc_info=True,
)
return True
except Exception as release_error:
logger.critical(
"Billing reservation release failed",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(release_error),
},
exc_info=True,
)
return True
async def handle_streaming_chat_completion(
self,
response: httpx.Response,
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
legacy_completion: bool = False,
) -> StreamingResponse:
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
Args:
response: Streaming response from upstream
key: API key for the authenticated user
max_cost_for_model: Maximum cost deducted upfront for the model
Returns:
StreamingResponse with cost data injected at the end
"""
if reservation_snapshot is None:
async with create_session() as snapshot_session:
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
if snapshot_key is None:
raise RuntimeError("Billing key disappeared before streaming")
reservation_snapshot = await get_reservation_snapshot(
snapshot_key, snapshot_session
)
usage_estimator = MissingUsageEstimator(request_body, model_obj)
logger.debug(
"Processing streaming chat completion",
extra={
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"response_status": response.status_code,
},
)
usage_finalized = False
last_model_seen: str | None = None
provider_seen: str | None = None
async def finalize_db_only() -> None:
nonlocal usage_finalized
if usage_finalized:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
return
try:
await adjust_payment_for_tokens(
fresh_key,
usage_estimator.response_data(last_model_seen),
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except Exception:
logger.exception(
"Fallback stream billing finalization failed; releasing reservation",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key, new_session, reservation_snapshot
)
)
except Exception:
logger.exception(
"Fallback stream billing recovery could not access the database",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
stream_finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
)
)
async def stream_with_cost(
max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]:
nonlocal usage_finalized, last_model_seen
usage_chunk_data: dict | None = None
done_seen: bool = False
stream_id: str | None = None
def _process_event(
raw_event: bytes, final: bool = False
) -> Iterator[bytes]:
"""Process one complete SSE event block (lines up to a blank line).
Handles arbitrary upstream framing across every supported
provider:
* ``data:`` lines are gathered and concatenated per the SSE
spec, so a payload split across network chunks is reassembled
before parsing.
* Comment/keepalive lines (those beginning with ``:`` such as
OpenRouter's ``: OPENROUTER PROCESSING``) are dropped. They
carry no JSON and forwarding them downstream breaks naive SSE
clients; the keepalive only matters for the upstream hop.
* Other SSE fields (``event:``/``id:``/``retry:``) are preserved
and kept attached to the event's ``data:`` line, which the
OpenAI Responses API and Anthropic-style streams rely on.
* ``[DONE]`` is swallowed so it can be re-emitted exactly once at
end of stream.
"""
nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id
nonlocal provider_seen
event = raw_event.strip(b"\r\n")
if not event:
return
field_lines: list[bytes] = []
data_lines: list[bytes] = []
for line in event.split(b"\n"):
line = line.rstrip(b"\r")
if line.startswith(b"data:"):
# Strip the field name and a single optional leading space.
data_lines.append(line[len(b"data:") :].lstrip(b" "))
elif line.startswith(b":"):
# SSE comment / keepalive - drop.
continue
elif line:
# Other SSE field (event:/id:/retry:) - preserve in order.
field_lines.append(line)
if not data_lines:
return
data = b"\n".join(data_lines)
if not data.strip():
return
# Re-emit preserved SSE fields immediately before the data line so
# event/data framing stays intact (single trailing newline each;
# the blank-line terminator is appended to the data line below).
prefix = b"".join(fl + b"\n" for fl in field_lines)
if data.strip() == b"[DONE]":
done_seen = True
return
obj = json_codec.loads(data)
if isinstance(obj, dict):
usage_estimator.observe(obj)
provider_seen = self._stamp_streamed_provider(obj, provider_seen)
if obj.get("model"):
last_model_seen = str(obj.get("model"))
if requested_model:
obj["model"] = requested_model
if (
"id" not in obj
or not isinstance(obj["id"], str)
or obj["id"] == "existing-id"
):
if stream_id is None:
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
stream_id = f"{id_prefix}-{uuid.uuid4()}"
obj["id"] = stream_id
else:
stream_id = obj["id"]
if isinstance(obj.get("usage"), dict):
# Capture usage for end-of-stream cost reconciliation.
# Some models (e.g. Gemini thinking models over the
# OpenAI-compat endpoint) attach ``usage`` to the SAME
# chunk that carries the final content/finish_reason
# rather than sending a separate ``choices: []`` usage
# chunk. Only swallow the chunk when it is a pure usage
# chunk (no choices); otherwise the content would be
# silently dropped and the client would receive no
# assistant message at all.
if obj.get("choices"):
# Capture usage (with model) for the cost trailer,
# but with choices stripped so the trailer never
# re-emits this chunk's content.
usage_chunk_data = {
k: v for k, v in obj.items() if k != "choices"
}
usage_chunk_data["choices"] = []
# Forward the content now, without usage, so token
# usage is reported exactly once (in the trailer).
forward = {k: v for k, v in obj.items() if k != "usage"}
yield (
prefix + b"data: " + json_codec.dumps(forward) + b"\n\n"
)
return
usage_chunk_data = obj
return
yield prefix + b"data: " + json_codec.dumps(obj) + b"\n\n"
else:
if final:
# Final flush of a truncated tail: the upstream closed
# mid-event, so ``data`` is incomplete JSON. Emitting it
# as a ``data:`` frame would hand the client invalid
# JSON (the "unexpected token" parse error). Drop it.
return
# Non-JSON data payload (partial fragment already reassembled
# by buffering, or a provider control string). Re-prefix each
# line so multi-line ``data`` stays valid SSE framing - a bare
# second line would otherwise reach the client without its
# ``data:`` field and break naive parsers.
body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n"))
yield prefix + body + b"\n"
try:
# Buffer bytes across network chunks and dispatch only on the SSE
# event delimiter (a blank line). ``aiter_bytes`` yields arbitrary
# byte boundaries, so a single event's JSON can span chunks and
# multiple events can arrive together; buffering makes parsing
# boundary-independent for every provider.
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
# Flush any trailing event that lacked a final blank line.
buffer = splitter.flush()
if buffer.strip():
for out in _process_event(buffer, final=True):
yield out
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
cost_data: dict
try:
adjustment_input = usage_estimator.billing_data(
usage_chunk_data, last_model_seen
)
cost_data = await adjust_payment_for_tokens(
fresh_key,
adjustment_input,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except BaseException as e:
logger.critical(
"Error during usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
# Release is a terminal billing state. Do not enqueue
# finalize_db_only from the generator's finally block
# and charge this request later.
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key,
session,
reservation_snapshot,
)
)
raise
if usage_chunk_data is None:
if stream_id is None:
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
stream_id = f"{id_prefix}-{uuid.uuid4()}"
usage_chunk_data = {
"id": stream_id,
"object": "text_completion"
if legacy_completion
else "chat.completion.chunk",
"model": last_model_seen or "unknown",
"provider": provider_seen,
"choices": [],
"usage": {
"prompt_tokens": cost_data.get("input_tokens", 0),
"completion_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
}
try:
self.inject_cost_metadata(
usage_chunk_data, cost_data, fresh_key
)
except Exception:
logger.exception(
"Failed to inject cost metadata into streaming chunk",
extra={
"key_hash": key.hashed_key[:8] + "...",
},
)
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error:
logger.warning(
"Upstream stream ended before the response was complete",
extra={
"error": str(stream_error),
"key_hash": key.hashed_key[:8] + "...",
},
)
except Exception as stream_error:
logger.warning(
"Streaming interrupted; finalizing before closing upstream",
extra={
"error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
finally:
await stream_finalizer.run()
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return ClosingStreamingResponse(
stream_with_cost(max_cost_for_model),
finalizer=stream_finalizer,
status_code=response.status_code,
headers=response_headers,
)
async def handle_non_streaming_chat_completion(
self,
response: httpx.Response,
key: ApiKey,
session: AsyncSession,
deducted_max_cost: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
legacy_completion: bool = False,
) -> Response:
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
Args:
response: Response from upstream
key: API key for the authenticated user
session: Database session for updating balance
deducted_max_cost: Maximum cost deducted upfront
Returns:
Response with cost data added to JSON body
"""
logger.debug(
"Processing non-streaming chat completion",
extra={
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"response_status": response.status_code,
},
)
content: bytes | None = None
try:
content = await response.aread()
response_json = json.loads(content)
self._apply_provider_field(response_json)
logger.debug(
"Parsed response JSON",
extra={
"key_hash": key.hashed_key[:8] + "...",
"model": response_json.get("model", "unknown"),
"has_usage": "usage" in response_json,
},
)
if requested_model:
response_json["model"] = requested_model
if "id" not in response_json or not isinstance(response_json["id"], str):
prefix = "cmpl" if legacy_completion else "chatcmpl"
response_json["id"] = f"{prefix}-{uuid.uuid4()}"
usage = response_json.get("usage")
if not isinstance(usage, dict) or not usage:
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
response_json["usage"] = usage_estimator.openai_response_data(
response_json.get("model")
)["usage"]
cost_data = await adjust_payment_for_tokens(
key,
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
await session.refresh(key)
remaining_balance_msats = key.balance
# Merge the shared cost contract into usage for SDKs and OpenCode.
if "usage" in response_json:
_inject_cost_into_usage(response_json, cost_data)
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
self._fold_cache_into_input_tokens(response_json["usage"])
published_cost = _published_cost(cost_data)
published_cost["sats_cost"] = published_cost["total_msats"] // 1000
published_cost["remaining_balance_msats"] = remaining_balance_msats
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": published_cost.copy()}
response_json["cost"] = published_cost
logger.debug(
"Payment adjustment completed for non-streaming",
extra={
"key_hash": key.hashed_key[:8] + "...",
"cost_data": cost_data,
"model": response_json.get("model", "unknown"),
"balance_after_adjustment": key.balance,
},
)
allowed_headers = {
"content-type",
"cache-control",
"date",
"vary",
"access-control-allow-origin",
"access-control-allow-methods",
"access-control-allow-headers",
"access-control-allow-credentials",
"access-control-expose-headers",
"access-control-max-age",
}
response_headers = {
k: v
for k, v in response.headers.items()
if k.lower() in allowed_headers
}
_inject_cost_response_headers(response_headers, cost_data)
if requested_model:
response_json["model"] = requested_model
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError as e:
logger.error(
"Failed to parse JSON from upstream response",
extra={
"error": str(e),
"key_hash": key.hashed_key[:8] + "...",
"content_preview": content[:200].decode(errors="ignore")
if content
else "empty",
},
)
raise
except Exception as e:
logger.error(
"Error processing non-streaming chat completion",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
async def handle_streaming_responses_completion(
self,
response: httpx.Response,
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
Args:
response: Streaming response from upstream
key: API key for the authenticated user
max_cost_for_model: Maximum cost deducted upfront for the model
Returns:
StreamingResponse with cost data injected at the end
"""
usage_estimator = MissingUsageEstimator(request_body, model_obj)
logger.debug(
"Processing streaming Responses API completion",
extra={
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"response_status": response.status_code,
},
)
usage_finalized = False
last_model_seen: str | None = None
provider_seen: str | None = None
async def finalize_db_only() -> None:
nonlocal usage_finalized
if usage_finalized:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
return
try:
await adjust_payment_for_tokens(
fresh_key,
usage_estimator.response_data(last_model_seen),
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except Exception:
logger.exception(
"Fallback Responses billing finalization failed; releasing reservation",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key, new_session, reservation_snapshot
)
)
except Exception:
logger.exception(
"Fallback Responses billing recovery could not access the database",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
stream_finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
)
)
async def stream_with_responses_cost(
max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]:
nonlocal usage_finalized, last_model_seen
reasoning_tokens: int = 0
usage_chunk_data: dict | None = None
done_seen: bool = False
def _process_event(
raw_event: bytes, final: bool = False
) -> Iterator[bytes]:
"""Process one complete SSE event block for the Responses API.
Buffers full events (delimited by a blank line) so parsing is
boundary-independent, gathers ``data:`` lines, drops comment/
keepalive lines (e.g. OpenRouter's ``: OPENROUTER PROCESSING``),
and preserves ``event:``/``id:`` fields attached to their data
line so Responses API event framing stays intact.
"""
nonlocal last_model_seen, usage_chunk_data, done_seen, provider_seen
nonlocal reasoning_tokens
event = raw_event.strip(b"\r\n")
if not event:
return
field_lines: list[bytes] = []
data_lines: list[bytes] = []
for line in event.split(b"\n"):
line = line.rstrip(b"\r")
if line.startswith(b"data:"):
data_lines.append(line[len(b"data:") :].lstrip(b" "))
elif line.startswith(b":"):
# SSE comment / keepalive - drop.
continue
elif line:
# Preserve event:/id:/retry: (Responses API event names).
field_lines.append(line)
if not data_lines:
return
data = b"\n".join(data_lines)
if not data.strip():
return
prefix = b"".join(fl + b"\n" for fl in field_lines)
if data.strip() == b"[DONE]":
done_seen = True
return
obj = json_codec.loads(data)
if isinstance(obj, dict):
provider_seen = self._stamp_streamed_provider(obj, provider_seen)
if obj.get("model"):
last_model_seen = str(obj.get("model"))
if requested_model:
obj["model"] = requested_model
# Track reasoning tokens for Responses API
if usage := obj.get("usage", {}):
if isinstance(usage, dict) and "reasoning_tokens" in usage:
reasoning_tokens += usage.get("reasoning_tokens", 0)
# Responses API usage is in response.completed/incomplete events
chunk_type = obj.get("type", "")
if chunk_type in (
"response.completed",
"response.incomplete",
):
usage_chunk_data = obj
if not usage_estimator.output_text:
usage_estimator.observe(obj)
return
usage_estimator.observe(obj)
yield prefix + b"data: " + json_codec.dumps(obj) + b"\n\n"
else:
if final:
# Final flush of a truncated tail: upstream closed
# mid-event, so ``data`` is incomplete JSON. Dropping it
# avoids handing the client an invalid ``data:`` frame.
return
# Re-prefix each line so multi-line ``data`` stays valid SSE
# framing for the client.
body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n"))
yield prefix + body + b"\n"
try:
# Buffer across network chunks; dispatch only on the SSE event
# delimiter so parsing is independent of byte boundaries.
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
buffer = splitter.flush()
if buffer.strip():
for out in _process_event(buffer, final=True):
yield out
# Always emit a cost-bearing data chunk
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
cost_data: dict
try:
adjustment_input = usage_estimator.billing_data(
usage_chunk_data, last_model_seen
)
cost_data = await adjust_payment_for_tokens(
fresh_key,
adjustment_input,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except BaseException as e:
logger.critical(
"Error during Responses API usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key,
session,
reservation_snapshot,
)
)
raise
if usage_chunk_data is None:
usage_chunk_data = {
"type": "response.completed",
"provider": provider_seen,
"response": {
"model": last_model_seen or "unknown",
"usage": {
"input_tokens": cost_data.get(
"input_tokens", 0
),
"output_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
},
"usage": {
"input_tokens": cost_data.get("input_tokens", 0),
"output_tokens": cost_data.get("output_tokens", 0),
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
}
try:
self.inject_cost_metadata(
usage_chunk_data, cost_data, fresh_key
)
except Exception:
logger.exception(
"Failed to inject cost metadata into Responses streaming chunk",
extra={
"key_hash": key.hashed_key[:8] + "...",
},
)
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error:
logger.warning(
"Upstream Responses API stream ended before the response was complete",
extra={
"error": str(stream_error),
"key_hash": key.hashed_key[:8] + "...",
},
)
except Exception as stream_error:
logger.warning(
"Responses API streaming interrupted; finalizing before closing upstream",
extra={
"error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
finally:
await stream_finalizer.run()
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return ClosingStreamingResponse(
stream_with_responses_cost(max_cost_for_model),
finalizer=stream_finalizer,
status_code=response.status_code,
headers=response_headers,
)
async def handle_non_streaming_responses_completion(
self,
response: httpx.Response,
key: ApiKey,
session: AsyncSession,
deducted_max_cost: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> Response:
"""Handle non-streaming Responses API responses with token usage tracking and cost adjustment.
Args:
response: Response from upstream
key: API key for the authenticated user
session: Database session for updating balance
deducted_max_cost: Maximum cost deducted upfront
Returns:
Response with cost data added to JSON body
"""
logger.debug(
"Processing non-streaming Responses API completion",
extra={
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"response_status": response.status_code,
},
)
content: bytes | None = None
try:
content = await response.aread()
response_json = json.loads(content)
self._apply_provider_field(response_json)
logger.debug(
"Parsed Responses API response JSON",
extra={
"key_hash": key.hashed_key[:8] + "...",
"model": response_json.get("model", "unknown"),
"has_usage": "usage" in response_json,
"has_reasoning_tokens": "usage" in response_json
and isinstance(response_json.get("usage"), dict)
and "reasoning_tokens" in response_json["usage"],
},
)
if not isinstance(response_json.get("usage"), dict):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
response_json["usage"] = usage_estimator.response_data(
response_json.get("model")
)["usage"]
if requested_model:
response_json["model"] = requested_model
if "id" not in response_json or not isinstance(response_json["id"], str):
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
cost_data = await adjust_payment_for_tokens(
key,
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
await session.refresh(key)
remaining_balance_msats = key.balance
# Merge the shared cost contract into usage for SDKs and OpenCode.
if "usage" in response_json:
_inject_cost_into_usage(response_json, cost_data)
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
self._fold_cache_into_input_tokens(response_json["usage"])
published_cost = _published_cost(cost_data)
published_cost["sats_cost"] = published_cost["total_msats"] // 1000
published_cost["remaining_balance_msats"] = remaining_balance_msats
response_json["metadata"] = response_json.get("metadata", {})
response_json["metadata"]["routstr"] = {"cost": published_cost.copy()}
response_json["cost"] = published_cost
logger.debug(
"Payment adjustment completed for non-streaming Responses API",
extra={
"key_hash": key.hashed_key[:8] + "...",
"cost_data": cost_data,
"model": response_json.get("model", "unknown"),
"balance_after_adjustment": key.balance,
},
)
allowed_headers = {
"content-type",
"cache-control",
"date",
"vary",
"access-control-allow-origin",
"access-control-allow-methods",
"access-control-allow-headers",
"access-control-allow-credentials",
"access-control-expose-headers",
"access-control-max-age",
}
response_headers = {
k: v
for k, v in response.headers.items()
if k.lower() in allowed_headers
}
_inject_cost_response_headers(response_headers, cost_data)
if requested_model:
response_json["model"] = requested_model
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError as e:
logger.error(
"Failed to parse JSON from upstream Responses API response",
extra={
"error": str(e),
"key_hash": key.hashed_key[:8] + "...",
"content_preview": content[:200].decode(errors="ignore")
if content
else "empty",
},
)
raise
except Exception as e:
logger.error(
"Error processing non-streaming Responses API completion",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
async def _finalize_generic_streaming_payment(
self,
key_hash: str,
max_cost: int,
path: str,
model_obj: Model | None,
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
) -> None:
"""Finalize payment for a generic streaming request."""
async with create_session() as session:
key = await session.get(ApiKey, key_hash)
if not key:
logger.warning(
"Key not found during generic streaming payment finalization",
extra={"key_hash": key_hash[:8] + "..."},
)
return
try:
# Generic opaque streams have no request/response token seam.
# Missing usage therefore releases the reservation; the hold is
# never treated as evidence of consumption.
await adjust_payment_for_tokens(
key,
{"model": "unknown", "usage": None},
session,
max_cost,
model_obj=model_obj,
provider_fee=provider_fee,
reservation_snapshot=reservation_snapshot,
)
logger.debug(
"Finalized generic streaming payment",
extra={
"path": path,
"key_hash": key_hash[:8] + "...",
},
)
except Exception as e:
logger.error(
"Error finalizing generic streaming payment",
extra={
"error": str(e),
"key_hash": key_hash[:8] + "...",
"path": path,
},
)
async def _stream_generic_with_settlement(
self,
response: httpx.Response,
key_hash: str,
max_cost: int,
path: str,
model_obj: Model | None,
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
finalizer: PersistentStreamFinalizer | None = None,
) -> AsyncGenerator[bytes, None]:
"""Relay an opaque stream and settle it even if the caller disconnects."""
if finalizer is None:
finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(
lambda: self._finalize_generic_streaming_payment(
key_hash,
max_cost,
path,
model_obj,
provider_fee,
reservation_snapshot,
),
response,
)
)
try:
async for chunk in response.aiter_bytes():
yield chunk
finally:
await finalizer.run()
def _generic_streaming_response(
self,
response: httpx.Response,
key_hash: str,
max_cost: int,
path: str,
model_obj: Model | None,
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
) -> ClosingStreamingResponse:
finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(
lambda: self._finalize_generic_streaming_payment(
key_hash,
max_cost,
path,
model_obj,
provider_fee,
reservation_snapshot,
),
response,
)
)
stream = self._stream_generic_with_settlement(
response,
key_hash,
max_cost,
path,
model_obj,
provider_fee,
reservation_snapshot,
finalizer,
)
return ClosingStreamingResponse(
stream,
finalizer=finalizer,
status_code=response.status_code,
headers=dict(response.headers),
)
async def handle_streaming_messages_completion(
self,
response: httpx.Response,
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_finalized = False
last_model_seen: str | None = None
provider_seen: str | None = None
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
if usage_finalized:
return None
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
usage_finalized = True
return None
try:
cost_data = await adjust_payment_for_tokens(
fresh_key,
usage_estimator.response_data(last_model_seen),
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
except BaseException as e:
logger.critical(
"Error during Messages API usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
usage_finalized = await self._release_failed_streaming_reservation(
fresh_key,
new_session,
reservation_snapshot,
)
raise
async def finalize_db_only() -> None:
if not usage_finalized:
await finalize_without_usage()
stream_finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(finalize_db_only, response)
)
async def stream_with_cost(
max_cost_for_model: int,
) -> AsyncGenerator[bytes, None]:
nonlocal usage_finalized, last_model_seen, provider_seen
stored_chunks: list[bytes] = []
input_tokens: int = 0
output_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_input_tokens: int = 0
total_cost: float = 0.0
input_cost: float = 0.0
output_cost: float = 0.0
def _coerce_usd(value: object) -> float:
if value is None or isinstance(value, bool):
return 0.0
if not isinstance(value, (int, float, str)):
return 0.0
try:
return max(0.0, float(value))
except (TypeError, ValueError):
return 0.0
def _absorb_usd(usage_or_root: dict) -> None:
nonlocal total_cost, input_cost, output_cost
cd = usage_or_root.get("cost_details")
if isinstance(cd, dict):
total_cost = max(
total_cost,
_coerce_usd(cd.get("total_cost")),
)
input_cost = max(
input_cost,
_coerce_usd(cd.get("input_cost")),
)
output_cost = max(
output_cost,
_coerce_usd(cd.get("output_cost")),
)
for field in ("total_cost", "cost"):
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
try:
async for chunk in response.aiter_bytes():
stored_chunks.append(chunk)
try:
decoded_chunk = chunk.decode("utf-8", errors="ignore")
modified_lines = []
changed = False
for line in decoded_chunk.split("\n"):
if line.startswith("data: "):
try:
data = json.loads(line[6:])
if isinstance(data, dict):
usage_estimator.observe(data)
msg = data.get("message", {})
if msg and msg.get("model"):
last_model_seen = str(msg.get("model"))
provider_added = "provider" not in data
provider_seen = self._stamp_streamed_provider(
data, provider_seen
)
if requested_model:
# Apply requested_model override
model_updated = False
if msg:
msg["model"] = requested_model
model_updated = True
if data.get("model"):
data["model"] = requested_model
model_updated = True
if model_updated or provider_added:
line = "data: " + json.dumps(data)
changed = True
elif provider_added:
line = "data: " + json.dumps(data)
changed = True
if usage := msg.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
# Anthropic's `message_start.usage`
# carries the cumulative cache
# snapshot for the prompt — pick
# the max() so subsequent
# `message_delta.usage` events
# (which only restate the same
# numbers) don't double-count.
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usd(usage)
if usage := data.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usd(usage)
# Some upstreams attach cost fields at
# the event root rather than nested
# under `usage`.
_absorb_usd(data)
except json.JSONDecodeError:
pass
modified_lines.append(line)
if changed:
yield "\n".join(modified_lines).encode("utf-8")
else:
yield chunk
except Exception:
yield chunk
usage_data = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
}
messages_dispatch.embed_usd_costs(
usage_data,
total_cost,
input_cost,
output_cost,
)
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
combined_data = {
"model": last_model_seen or "unknown",
"provider": provider_seen,
"usage": usage_data,
}
cost_data = await adjust_payment_for_tokens(
fresh_key,
combined_data,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(
combined_data, cost_data, fresh_key
)
usage_finalized = True
# Emit the full combined_data as the cost
yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode()
except BaseException as e:
logger.critical(
"Error during Messages API usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key,
new_session,
reservation_snapshot,
)
)
raise
if not usage_finalized:
maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None:
yield maybe_cost_event
except httpx.ReadError:
if not usage_finalized:
await finalize_without_usage()
# Upstream dropped the connection mid-stream; response already started, swallow silently
except Exception:
if not usage_finalized:
await finalize_without_usage()
raise
finally:
await stream_finalizer.run()
response_headers = dict(response.headers)
response_headers.pop("content-encoding", None)
response_headers.pop("content-length", None)
return ClosingStreamingResponse(
stream_with_cost(max_cost_for_model),
finalizer=stream_finalizer,
status_code=response.status_code,
headers=response_headers,
)
async def handle_non_streaming_messages_completion(
self,
response: httpx.Response,
key: ApiKey,
session: AsyncSession,
deducted_max_cost: int,
path: str,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> Response:
try:
content = await response.aread()
response_json = json.loads(content)
if requested_model:
if "model" in response_json:
response_json["model"] = requested_model
if (
"message" in response_json
and isinstance(response_json["message"], dict)
and "model" in response_json["message"]
):
response_json["message"]["model"] = requested_model
if path.endswith("count_tokens") and "usage" not in response_json:
input_tokens = response_json.get("input_tokens", 0)
response_json["usage"] = {"input_tokens": input_tokens}
elif not isinstance(response_json.get("usage"), dict):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
response_json["usage"] = usage_estimator.response_data(
response_json.get("model")
)["usage"]
cost_data = await adjust_payment_for_tokens(
key,
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(response_json, cost_data, key)
allowed_headers = {
"content-type",
"cache-control",
"date",
"vary",
"access-control-allow-origin",
"access-control-allow-methods",
"access-control-allow-headers",
"access-control-allow-credentials",
"access-control-expose-headers",
"access-control-max-age",
}
response_headers = {
k: v
for k, v in response.headers.items()
if k.lower() in allowed_headers
}
# Inject the same cost headers used by every paid response path.
_inject_cost_response_headers(response_headers, cost_data)
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except Exception:
raise
# ------------------------------------------------------------------
# Litellm /v1/messages dispatch (thin wrappers)
#
# The actual translation logic lives in ``messages_dispatch``. These
# method shims exist so subclasses and tests can keep the original
# provider-bound API.
# ------------------------------------------------------------------
_coerce_litellm_payload = staticmethod(messages_dispatch.coerce_litellm_payload)
_parse_sse_blocks = staticmethod(messages_dispatch.parse_sse_blocks)
_events_from_chunk = staticmethod(messages_dispatch.events_from_chunk)
async def _aggregate_anthropic_events_to_message(
self, iterator: AsyncIterator[Any]
) -> dict:
return await messages_dispatch.aggregate_anthropic_events_to_message(iterator)
def adapt_messages_request(self, body: dict, model_obj: Model) -> str:
"""Rewrite an allowlisted /v1/messages body for this upstream.
Returns a suffix appended to the upstream model name, empty when the
provider needs none. Subclasses override this to express an Anthropic
feature the upstream spells differently; the base forwards the body
untouched.
"""
return ""
async def _dispatch_anthropic_messages(
self,
request_body: bytes | None,
model_obj: Model,
*,
log_extra: dict[str, Any] | None = None,
) -> tuple[bool, Any, str | None]:
return await messages_dispatch.dispatch_anthropic_messages(
request_body=request_body,
model_obj=model_obj,
base_url=self.base_url,
api_key=self.api_key,
provider_prefix=self.get_litellm_provider_prefix(),
transform_model_name=self.transform_model_name,
adapt_request=lambda body: self.adapt_messages_request(body, model_obj),
log_extra=log_extra,
)
async def _forward_messages_via_litellm(
self,
request_body: bytes | None,
key: ApiKey,
session: AsyncSession,
max_cost_for_model: int,
model_obj: Model,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response | StreamingResponse:
"""Translate /v1/messages to upstream chat/completions via litellm.
Used when the upstream provider does not natively serve Anthropic
Messages (i.e. supports_anthropic_messages is False). Cost
tracking and metadata injection mirror the native messages path.
"""
stream, result, requested_model = await self._dispatch_anthropic_messages(
request_body,
model_obj,
log_extra={"key_hash": key.hashed_key[:8] + "..."},
)
if stream:
return self._stream_litellm_messages(
cast(AsyncIterator[Any], result),
key,
max_cost_for_model,
requested_model,
model_obj,
reservation_snapshot,
request_body,
)
response_json = messages_dispatch.coerce_litellm_payload(result)
if requested_model and "model" in response_json:
response_json["model"] = requested_model
if not isinstance(response_json.get("usage"), dict):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
response_json["usage"] = usage_estimator.response_data(
response_json.get("model")
)["usage"]
cost_data = await adjust_payment_for_tokens(
key,
response_json,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(response_json, cost_data, key)
# Inject the same cost headers used by every paid response path.
response_headers: dict[str, str] = {}
_inject_cost_response_headers(response_headers, cost_data)
return Response(
content=json.dumps(response_json).encode(),
status_code=200,
headers=response_headers,
media_type="application/json",
)
async def _forward_x_cashu_messages_via_litellm(
self,
request_body: bytes,
amount: int,
unit: str,
max_cost_for_model: int,
model_obj: Model,
mint: str | None = None,
request_id: str | None = None,
) -> Response | StreamingResponse:
"""Dispatch /v1/messages via litellm for x-cashu payments.
Computes cost from upstream usage, refunds the unspent balance via
an X-Cashu response header, and returns the Anthropic-shaped body.
"""
stream, result, requested_model = await self._dispatch_anthropic_messages(
request_body,
model_obj,
log_extra={"payment_unit": unit, "payment_amount": amount},
)
if stream:
return await self._stream_x_cashu_litellm_messages(
cast(AsyncIterator[Any], result),
amount,
unit,
max_cost_for_model,
requested_model,
mint,
request_id,
model_obj,
)
response_json = messages_dispatch.coerce_litellm_payload(result)
self._apply_provider_field(response_json)
if requested_model and "model" in response_json:
response_json["model"] = requested_model
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if (
cost_data
and "usage" in response_json
and isinstance(response_json["usage"], dict)
):
_inject_cost_into_usage(response_json, cost_data)
self._fold_cache_into_input_tokens(response_json["usage"])
response_headers: dict[str, str] = {}
if cost_data:
_inject_cost_response_headers(response_headers, cost_data)
refund_amount = messages_dispatch.compute_refund(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for non-streaming /v1/messages via litellm",
extra={
"refund_amount": refund_amount,
"unit": unit,
"model": response_json.get("model", "unknown"),
},
)
return Response(
content=json.dumps(response_json).encode(),
status_code=200,
headers=response_headers,
media_type="application/json",
)
_compute_refund = staticmethod(messages_dispatch.compute_refund)
def _stream_litellm_messages(
self,
iterator: AsyncIterator[Any],
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
"""Re-emit a litellm Anthropic-event iterator as live SSE bytes
with cost reconciliation appended at end of stream."""
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_finalized = False
last_model_seen: str | None = None
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
if usage_finalized:
return None
logger.warning(
"Finalizing /v1/messages stream with locally estimated "
"usage because the upstream omitted `usage` from SSE. "
"Check that the upstream emits a final usage chunk; the "
"reservation ceiling will not be used as the charge.",
extra={
"key_hash": key.hashed_key[:8] + "...",
"model": last_model_seen or "unknown",
"provider": self.provider_type or self.base_url,
"max_cost_msats": max_cost_for_model,
},
)
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
usage_finalized = True
return None
try:
cost_data = await adjust_payment_for_tokens(
fresh_key,
usage_estimator.response_data(last_model_seen),
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
return (
f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n"
).encode()
except BaseException as e:
logger.critical(
"Error during LiteLLM Messages usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
usage_finalized = await self._release_failed_streaming_reservation(
fresh_key,
new_session,
reservation_snapshot,
)
raise
async def finalize_stream() -> None:
try:
if not usage_finalized:
await finalize_without_usage()
finally:
await aclose_if_needed(iterator)
stream_finalizer = PersistentStreamFinalizer(finalize_stream)
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
nonlocal usage_finalized, last_model_seen
input_tokens = 0
output_tokens = 0
cache_read_input_tokens = 0
cache_creation_input_tokens = 0
total_cost = 0.0
input_cost = 0.0
output_cost = 0.0
try:
async for annotated in messages_dispatch.stream_annotated_events(
iterator, requested_model
):
usage_estimator.observe(annotated.event)
if annotated.model:
last_model_seen = annotated.model
# Anthropic SSE reports usage cumulatively across
# message_start + message_delta — take the max snapshot
# rather than summing, otherwise input tokens
# double-count.
input_tokens = max(input_tokens, annotated.input_tokens)
output_tokens = max(output_tokens, annotated.output_tokens)
cache_read_input_tokens = max(
cache_read_input_tokens,
annotated.cache_read_input_tokens,
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
annotated.cache_creation_input_tokens,
)
total_cost = max(total_cost, annotated.total_cost)
input_cost = max(input_cost, annotated.input_cost)
output_cost = max(output_cost, annotated.output_cost)
yield annotated.sse_bytes
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
rebuilt_usage: dict = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": (
cache_read_input_tokens
),
"cache_creation_input_tokens": (
cache_creation_input_tokens
),
}
messages_dispatch.embed_usd_costs(
rebuilt_usage,
total_cost,
input_cost,
output_cost,
)
combined_data: dict = {
"model": last_model_seen or "unknown",
"usage": rebuilt_usage,
}
cost_data = await adjust_payment_for_tokens(
fresh_key,
combined_data,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(
combined_data, cost_data, fresh_key
)
usage_finalized = True
yield (
f"event: cost\ndata: "
f"{json.dumps({'cost': cost_data})}\n\n"
).encode()
except BaseException as e:
logger.critical(
"Error during LiteLLM Messages usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key,
new_session,
reservation_snapshot,
)
)
raise
if not usage_finalized:
cost_event = await finalize_without_usage()
if cost_event is not None:
yield cost_event
except Exception:
if not usage_finalized:
await finalize_without_usage()
raise
finally:
await stream_finalizer.run()
return ClosingStreamingResponse(
stream_with_cost(),
finalizer=stream_finalizer,
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
)
async def _stream_x_cashu_litellm_messages(
self,
iterator: AsyncIterator[Any],
amount: int,
unit: str,
max_cost_for_model: int,
requested_model: str | None,
mint: str | None,
request_id: str | None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Buffer a litellm stream end-to-end, compute cost, then replay.
Note this is **not** true streaming — the full event sequence is
accumulated into memory before a single byte is sent to the
client. The constraint is the ``X-Cashu`` refund token, which must
be set as a response *header* and therefore has to be known before
the response begins. The bearer-key path
(:meth:`_stream_litellm_messages`) avoids this by emitting cost as
a trailing ``event: cost`` SSE message; switching x-cashu to the
same trailing-event contract would let this path stream live, at
the cost of a wire-format change for clients that read ``X-Cashu``
from headers today.
"""
buffered: list[messages_dispatch.AnnotatedEvent] = []
last_model_seen: str | None = None
input_tokens = 0
output_tokens = 0
cache_read_input_tokens = 0
cache_creation_input_tokens = 0
total_cost = 0.0
input_cost = 0.0
output_cost = 0.0
async for annotated in messages_dispatch.stream_annotated_events(
iterator, requested_model
):
if annotated.model:
last_model_seen = annotated.model
# See _stream_litellm_messages for why this is max() not +=.
input_tokens = max(input_tokens, annotated.input_tokens)
output_tokens = max(output_tokens, annotated.output_tokens)
cache_read_input_tokens = max(
cache_read_input_tokens, annotated.cache_read_input_tokens
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
annotated.cache_creation_input_tokens,
)
total_cost = max(total_cost, annotated.total_cost)
input_cost = max(input_cost, annotated.input_cost)
output_cost = max(output_cost, annotated.output_cost)
buffered.append(annotated)
response_headers: dict[str, str] = {
"Cache-Control": "no-cache",
"Connection": "keep-alive",
}
if (
input_tokens == 0
and output_tokens == 0
and cache_read_input_tokens == 0
and cache_creation_input_tokens == 0
and total_cost == 0
):
logger.warning(
"x-cashu /v1/messages stream finished with no usage data "
"— refund cannot be computed and the client effectively "
"pays the full cashu amount. Likely cause: upstream "
"omitted `usage` from the SSE stream.",
extra={
"model": last_model_seen or "unknown",
"provider": self.provider_type or self.base_url,
"amount": amount,
"unit": unit,
},
)
cost_data: CostData | MaxCostData | None = None
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
rebuilt_usage: dict = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
}
messages_dispatch.embed_usd_costs(
rebuilt_usage, total_cost, input_cost, output_cost
)
response_data: dict = {
"model": last_model_seen or "unknown",
"usage": rebuilt_usage,
}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data:
refund_amount = messages_dispatch.compute_refund(
amount, unit, cost_data.total_msats
)
if refund_amount > 0:
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming /v1/messages via litellm",
extra={
"refund_amount": refund_amount,
"unit": unit,
"model": last_model_seen,
},
)
except Exception as exc:
logger.error(
"Error calculating cost for streaming /v1/messages",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"amount": amount,
"unit": unit,
},
)
if cost_data:
_inject_cost_response_headers(response_headers, cost_data)
for index, annotated in enumerate(buffered):
event = annotated.event
changed = False
message = event.get("message")
if isinstance(message, dict) and isinstance(message.get("usage"), dict):
_inject_cost_into_usage(message, cost_data)
changed = True
if isinstance(event.get("usage"), dict):
_inject_cost_into_usage(event, cost_data)
changed = True
if changed:
event_type = str(event.get("type") or "")
prefix = f"event: {event_type}\n" if event_type else ""
buffered[index] = annotated._replace(
sse_bytes=(f"{prefix}data: {json.dumps(event)}\n\n".encode())
)
async def replay() -> AsyncGenerator[bytes, None]:
for annotated in buffered:
yield annotated.sse_bytes
return ClosingStreamingResponse(
replay(),
media_type="text/event-stream",
headers=response_headers,
)
async def forward_request(
self,
request: Request,
path: str,
headers: dict,
request_body: bytes | None,
key: ApiKey,
max_cost_for_model: int,
session: AsyncSession,
model_obj: Model,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response | StreamingResponse:
"""Forward authenticated request to upstream service with cost tracking.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
request_body: Request body bytes, if any
key: API key for authenticated user
max_cost_for_model: Maximum cost deducted upfront
session: Database session for balance updates
Returns:
Response or StreamingResponse from upstream with cost tracking
"""
completion_path = _openai_completion_path(path)
path = self.normalize_request_path(path, model_obj)
if (
path.endswith("messages/count_tokens")
and not self.supports_anthropic_messages
):
return count_tokens_locally(request_body, model_obj)
if (
path.endswith("messages")
and not path.endswith("count_tokens")
and not self.supports_anthropic_messages
):
return await self._forward_messages_via_litellm(
request_body=request_body,
key=key,
session=session,
max_cost_for_model=max_cost_for_model,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
url = self.build_request_url(path, model_obj)
original_model_id = (
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
)
transformed_body = self.prepare_request_body(
request_body,
model_obj,
include_stream_usage=completion_path is not None,
)
logger.debug(
"Forwarding request to upstream",
extra={
"url": url,
"method": request.method,
"path": path,
"model": original_model_id or "unknown",
"provider": self.provider_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
response: httpx.Response | None = None
response_handoff = ResponseHandoff()
try:
client = acquire_upstream_http_client(url)
if transformed_body is not None:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=transformed_body,
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
else:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
response_handoff.acquire(response)
if response.status_code != 200:
if response.status_code >= 500:
try:
body_bytes = await response.aread()
except Exception:
body_bytes = b""
# Redact provider account identifiers before the body text
# reaches logs or the raised error.
body_preview = redact_org_ids(
body_bytes.decode("utf-8", errors="ignore").strip()[:500]
)
rate_limit = classify_rate_limit(
response.status_code,
body_preview,
dict(response.headers),
)
logger.error(
"Upstream %s returned %s for model=%s path=%s: %s",
self.provider_type,
response.status_code,
original_model_id or "unknown",
path,
body_preview or "<empty>",
extra={
"provider": self.provider_type,
"model": original_model_id or "unknown",
"status_code": response.status_code,
"error_code": rate_limit.code if rate_limit else None,
"reason_phrase": response.reason_phrase,
"path": path,
"body_preview": body_preview,
},
)
await response_handoff.close()
raise UpstreamError(
f"Upstream {self.provider_type} returned {response.status_code} "
f"for model {original_model_id or 'unknown'}: "
f"{body_preview[:200] or '<empty>'}",
status_code=response.status_code,
code=rate_limit.code if rate_limit else None,
details=rate_limit.as_details() if rate_limit else None,
from_upstream_response=True,
)
try:
mapped_error = await self.forward_upstream_error_response(
request, path, response, model_id=original_model_id
)
finally:
await response_handoff.close()
return mapped_error
if (
completion_path is not None
or path.endswith("embeddings")
or path.endswith("messages")
or path.endswith("messages/count_tokens")
or path.endswith("systemone")
or path.endswith("decisions")
):
if path.endswith("messages"):
client_wants_streaming = False
if request_body:
try:
request_data = json.loads(request_body)
client_wants_streaming = request_data.get("stream", False)
except json.JSONDecodeError:
pass
content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_messages_completion(
response,
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
)
response_handoff.handoff()
return result
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
response,
key,
session,
max_cost_for_model,
path,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
)
finally:
await response_handoff.close()
if path.endswith("messages/count_tokens"):
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
response,
key,
session,
max_cost_for_model,
path,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
)
finally:
await response_handoff.close()
if completion_path is not None:
client_wants_streaming = False
if request_body:
try:
request_data = json.loads(request_body)
client_wants_streaming = request_data.get("stream", False)
logger.debug(
"Chat completion request analysis",
extra={
"client_wants_streaming": client_wants_streaming,
"model": request_data.get("model", "unknown"),
"key_hash": key.hashed_key[:8] + "...",
},
)
except json.JSONDecodeError:
logger.warning(
"Failed to parse request body JSON for streaming detection"
)
content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming
logger.debug(
"Response type analysis",
extra={
"is_streaming": is_streaming,
"client_wants_streaming": client_wants_streaming,
"upstream_is_streaming": upstream_is_streaming,
"content_type": content_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_chat_completion(
response,
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
legacy_completion=completion_path == "completions",
)
response_handoff.handoff()
return result
# Handle both non-streaming chat completions and embeddings
if response.status_code == 200:
try:
return await self.handle_non_streaming_chat_completion(
response,
key,
session,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
legacy_completion=completion_path == "completions",
)
finally:
await response_handoff.close()
if reservation_snapshot is None:
reservation_snapshot = await get_reservation_snapshot(key, session)
logger.debug(
"Streaming non-chat response",
extra={
"path": path,
"status_code": response.status_code,
"key_hash": key.hashed_key[:8] + "...",
},
)
result = self._generic_streaming_response(
response,
key.hashed_key,
max_cost_for_model,
path,
model_obj,
self.provider_fee,
reservation_snapshot,
)
response_handoff.handoff()
return result
except UpstreamError:
await response_handoff.close()
raise
except httpx.RequestError as exc:
await response_handoff.close()
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),
"key_hash": key.hashed_key[:8] + "...",
},
)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
if isinstance(exc, httpx.PoolTimeout):
error_message = "Upstream connection pool is busy"
status_code = 503
elif isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
status_code = 502
elif isinstance(exc, httpx.TimeoutException):
error_message = "Upstream service request timed out"
status_code = 502
elif isinstance(exc, httpx.NetworkError):
error_message = "Network error while connecting to upstream service"
status_code = 502
else:
error_message = f"Error connecting to upstream service: {error_type}"
status_code = 502
raise UpstreamError(error_message, status_code=status_code)
except Exception as exc:
await response_handoff.close()
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),
"key_hash": key.hashed_key[:8] + "...",
"traceback": tb,
},
)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
raise UpstreamError(
"An unexpected server error occurred",
status_code=500,
scope=ERROR_SCOPE_NODE,
)
except BaseException:
await response_handoff.close(suppress_errors=True)
raise
supports_ehbp: bool = False
def get_confidential_inference_profile(
self,
) -> "ConfidentialInferenceProfile | None":
"""Return provider policy for encrypted/confidential inference forwarding."""
return None
def get_ehbp_forwarding_target(
self, path: str, model_obj: Model
) -> "EHBPForwardingTarget":
"""Return the EHBP forwarding target for this provider.
Providers must explicitly opt in by setting ``supports_ehbp = True``
and overriding this method. Most upstreams do not accept EHBP-encrypted
request bodies, so the base provider intentionally does not provide a
default endpoint.
"""
raise NotImplementedError(
f"Provider {self.provider_type} does not support EHBP forwarding"
)
async def forward_responses_request(
self,
request: Request,
path: str,
headers: dict,
request_body: bytes | None,
key: ApiKey,
max_cost_for_model: int,
session: AsyncSession,
model_obj: Model,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response | StreamingResponse:
"""Forward authenticated Responses API request to upstream service with cost tracking.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
request_body: Request body bytes, if any
key: API key for authenticated user
max_cost_for_model: Maximum cost deducted upfront
session: Database session for balance updates
model_obj: Model object for the request
Returns:
Response or StreamingResponse from upstream with cost tracking
"""
path = self.normalize_request_path(path, model_obj)
url = self.build_request_url(path, model_obj)
original_model_id = (
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
)
transformed_body = self.prepare_responses_request_body(request_body, model_obj)
logger.debug(
"Forwarding Responses API request to upstream",
extra={
"url": url,
"method": request.method,
"path": path,
"model": original_model_id or "unknown",
"provider": self.provider_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
response: httpx.Response | None = None
response_handoff = ResponseHandoff()
try:
client = acquire_upstream_http_client(url)
if transformed_body is not None:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=transformed_body,
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
else:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
response_handoff.acquire(response)
if response.status_code != 200:
if response.status_code >= 500:
try:
body_bytes = await response.aread()
except Exception:
body_bytes = b""
# Redact provider account identifiers before the body text
# reaches logs or the raised error.
body_preview = redact_org_ids(
body_bytes.decode("utf-8", errors="ignore").strip()[:500]
)
rate_limit = classify_rate_limit(
response.status_code,
body_preview,
dict(response.headers),
)
logger.error(
"Upstream %s returned %s for model=%s path=%s: %s",
self.provider_type,
response.status_code,
original_model_id or "unknown",
path,
body_preview or "<empty>",
extra={
"provider": self.provider_type,
"model": original_model_id or "unknown",
"status_code": response.status_code,
"error_code": rate_limit.code if rate_limit else None,
"path": path,
"body_preview": body_preview,
},
)
await response_handoff.close()
raise UpstreamError(
f"Upstream {self.provider_type} returned {response.status_code} "
f"for model {original_model_id or 'unknown'}: "
f"{body_preview[:200] or '<empty>'}",
status_code=response.status_code,
code=rate_limit.code if rate_limit else None,
details=rate_limit.as_details() if rate_limit else None,
from_upstream_response=True,
)
try:
mapped_error = await self.forward_upstream_error_response(
request, path, response, model_id=original_model_id
)
finally:
await response_handoff.close()
return mapped_error
if path.startswith("responses"):
content_type = response.headers.get("content-type", "")
is_streaming = "text/event-stream" in content_type
logger.debug(
"Responses API response type analysis",
extra={
"is_streaming": is_streaming,
"content_type": content_type,
"key_hash": key.hashed_key[:8] + "...",
},
)
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_responses_completion(
response,
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=transformed_body,
)
response_handoff.handoff()
return result
if response.status_code == 200:
try:
return await self.handle_non_streaming_responses_completion(
response,
key,
session,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=transformed_body,
)
finally:
await response_handoff.close()
if reservation_snapshot is None:
reservation_snapshot = await get_reservation_snapshot(key, session)
logger.debug(
"Streaming non-Responses API response",
extra={
"path": path,
"status_code": response.status_code,
"key_hash": key.hashed_key[:8] + "...",
},
)
result = self._generic_streaming_response(
response,
key.hashed_key,
max_cost_for_model,
path,
model_obj,
self.provider_fee,
reservation_snapshot,
)
response_handoff.handoff()
return result
except UpstreamError:
await response_handoff.close()
raise
except httpx.RequestError as exc:
await response_handoff.close()
error_type = type(exc).__name__
error_details = str(exc)
logger.error(
"HTTP request error to upstream Responses API",
extra={
"error_type": error_type,
"error_details": error_details,
"method": request.method,
"url": url,
"path": path,
"query_params": dict(request.query_params),
"key_hash": key.hashed_key[:8] + "...",
},
)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
if isinstance(exc, httpx.PoolTimeout):
error_message = "Upstream connection pool is busy"
status_code = 503
elif isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
status_code = 502
elif isinstance(exc, httpx.TimeoutException):
error_message = "Upstream service request timed out"
status_code = 502
elif isinstance(exc, httpx.NetworkError):
error_message = "Network error while connecting to upstream service"
status_code = 502
else:
error_message = f"Error connecting to upstream service: {error_type}"
status_code = 502
raise UpstreamError(error_message, status_code=status_code)
except Exception as exc:
await response_handoff.close()
tb = traceback.format_exc()
logger.error(
"Unexpected error in upstream Responses API forwarding",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"method": request.method,
"url": url,
"path": path,
"query_params": dict(request.query_params),
"key_hash": key.hashed_key[:8] + "...",
"traceback": tb,
},
)
# Don't revert here — proxy.py owns payment revert to avoid double-revert
raise UpstreamError(
"An unexpected server error occurred",
status_code=500,
scope=ERROR_SCOPE_NODE,
)
except BaseException:
await response_handoff.close(suppress_errors=True)
raise
async def forward_get_request(
self,
request: Request,
path: str,
headers: dict,
) -> Response | StreamingResponse:
"""Forward unauthenticated GET request to upstream service.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
Returns:
StreamingResponse from upstream
"""
path = self.normalize_request_path(path)
url = self.build_request_url(path)
logger.debug(
"Forwarding GET request to upstream",
extra={
"url": url,
"method": request.method,
"path": path,
"provider": self.provider_type,
},
)
response: httpx.Response | None = None
try:
client = acquire_upstream_http_client(url)
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=self.prepare_params(path, request.query_params),
),
)
logger.debug(
"GET request forwarded",
extra={
"path": path,
"status_code": response.status_code,
"provider": self.provider_type,
},
)
if response.status_code != 200:
return await self.forward_upstream_error_response(
request, path, response
)
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 UpstreamError:
raise
except httpx.PoolTimeout:
logger.warning(
"Upstream connection pool exhausted on GET",
extra={"path": path, "url": url, "provider": self.provider_type},
)
return create_error_response(
"service_unavailable",
"Upstream connection pool is busy",
503,
request=request,
)
except httpx.RequestError as exc:
logger.warning(
"Upstream request error on GET",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"path": path,
"url": url,
"provider": self.provider_type,
},
)
return create_error_response(
"upstream_error",
"Unable to reach upstream service",
502,
request=request,
)
except Exception as exc:
logger.error(
"Error forwarding GET request",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"method": request.method,
"url": url,
"path": path,
"query_params": dict(request.query_params),
"traceback": traceback.format_exc(),
},
)
return create_error_response(
"internal_error",
"An unexpected server error occurred",
500,
request=request,
)
finally:
await aclose_if_needed(response)
async def get_x_cashu_cost(
self,
response_data: dict,
max_cost_for_model: int,
model_obj: Model | None,
) -> MaxCostData | CostData | None:
"""Calculate cost for X-Cashu payment based on response data.
Args:
response_data: Response data containing model and usage information
max_cost_for_model: Maximum cost for the model
model_obj: The model that actually served the request; billed
directly instead of re-deriving pricing from the upstream's
echoed model string
Returns:
Cost data object (MaxCostData or CostData) or None if calculation fails
"""
model = response_data.get("model", None)
logger.debug(
"Calculating cost for response",
extra={"model": model, "has_usage": "usage" in response_data},
)
match await calculate_cost(
response_data,
max_cost_for_model,
model_obj,
self.provider_fee,
):
case MaxCostData() as cost:
logger.debug(
"Using max cost pricing",
extra={"model": model, "max_cost_msats": cost.total_msats},
)
return cost
case CostData() as cost:
logger.debug(
"Using token-based pricing",
extra={
"model": model,
"total_cost_msats": cost.total_msats,
"input_msats": cost.input_msats,
"output_msats": cost.output_msats,
},
)
return cost
case CostDataError() as error:
# Content was already served, so refund instead of raising.
logger.error(
"Cost calculation error, refunding the prepayment",
extra={
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
return MaxCostData(
base_msats=0, input_msats=0, output_msats=0, total_msats=0
)
return None
async def send_refund(
self,
amount: int,
unit: str,
mint: str | None = None,
request_id: str | None = None,
) -> str:
"""Create and send a refund token to the user.
Args:
amount: Refund amount
unit: Unit of the refund (sat or msat)
mint: Optional mint URL for the refund token
request_id: Optional HTTP request ID for tracking
Returns:
Refund token string
"""
logger.debug(
"Creating refund token",
extra={"amount": amount, "unit": unit, "mint": mint},
)
try:
# Token creation may swap proofs, so it is unsafe to retry.
refund_token = await send_token(amount, unit=unit, mint_url=mint)
except Exception as error:
logger.error(
"Failed to create refund token",
extra={
"error": str(error),
"error_type": type(error).__name__,
"amount": amount,
"unit": unit,
"mint": mint,
},
)
raise HTTPException(
status_code=401,
detail={
"error": {
"message": f"failed to create refund: {error}",
"type": "invalid_request_error",
"code": "send_token_failed",
}
},
) from error
logger.info(
"Refund token created successfully",
extra={
"amount": amount,
"unit": unit,
"mint": mint,
"token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
await store_cashu_transaction(
token=refund_token,
amount=amount,
unit=unit,
mint_url=token_mint_url(refund_token, mint),
typ="out",
request_id=request_id,
)
return refund_token
async def handle_x_cashu_streaming_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
"""Handle streaming response for X-Cashu payment, calculating refund if needed.
Args:
content_str: Response content as string
response: Original httpx response
amount: Payment amount received
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
Returns:
StreamingResponse with refund token in header if applicable
"""
logger.debug(
"Processing streaming response",
extra={
"amount": amount,
"unit": unit,
"content_lines": len(content_str.strip().split("\n")),
},
)
response_headers = dict(response.headers)
if "transfer-encoding" in response_headers:
del response_headers["transfer-encoding"]
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
usage_data = None
model = None
cost_data: CostData | MaxCostData | None = None
usage_estimator = MissingUsageEstimator(request_body, model_obj)
lines = content_str.strip().split("\n")
for line in lines:
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
# OpenAI format: usage and model at top level
if "usage" in data_json:
usage_data = data_json["usage"]
model = data_json.get("model") or model
elif "model" in data_json and not model:
model = data_json["model"]
# Anthropic format: model and input usage inside "message" key
if "message" in data_json:
msg = data_json["message"]
if not model and msg.get("model"):
model = msg["model"]
if msg.get("usage") and not usage_data:
usage_data = msg["usage"]
elif msg.get("usage") and usage_data:
# Merge: message_start has input_tokens, message_delta has output_tokens
merged = dict(usage_data)
for k, v in msg["usage"].items():
merged[k] = merged.get(k, 0) + v
usage_data = merged
except json.JSONDecodeError:
continue
usage_estimator.observe(data_json)
if not usage_data:
usage_data = usage_estimator.estimated_usage(model)
if usage_data:
logger.warning(
"No usage in streaming response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"estimated_usage": usage_data,
},
)
logger.debug(
"Calculating cost for streaming response",
extra={
"model": model,
"usage_data": usage_data,
"amount": amount,
"unit": unit,
},
)
response_data = {"usage": usage_data, "model": model or "unknown"}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="chat.streaming",
model=model,
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
if refund_amount > 0:
logger.debug(
"Processing refund for streaming response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
},
)
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
else:
logger.debug(
"No refund needed for streaming response",
extra={
"amount": amount,
"cost_msats": cost_data.total_msats,
"model": model,
},
)
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e:
logger.error(
"Error calculating cost for streaming response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"amount": amount,
"unit": unit,
},
)
provider_seen: str | None = None
for i, line in enumerate(lines):
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if not isinstance(data_json, dict):
continue
provider_before = data_json.get("provider")
provider_seen = self._stamp_streamed_provider(
data_json, provider_seen
)
changed = data_json.get("provider") != provider_before
if cost_data and "usage" in data_json and data_json["usage"]:
_inject_cost_into_usage(data_json, cost_data)
changed = True
if changed:
lines[i] = "data: " + json.dumps(data_json)
except json.JSONDecodeError:
pass
async def generate() -> AsyncGenerator[bytes, None]:
for line in lines:
yield (line + "\n").encode("utf-8")
return ClosingStreamingResponse(
generate(),
status_code=response.status_code,
headers=response_headers,
media_type="text/plain",
)
async def handle_x_cashu_non_streaming_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> Response:
"""Handle non-streaming response for X-Cashu payment, calculating refund if needed.
Args:
content_str: Response content as string
response: Original httpx response
amount: Payment amount received
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
Returns:
Response with refund token in header if applicable
"""
logger.debug(
"Processing non-streaming response",
extra={"amount": amount, "unit": unit, "content_length": len(content_str)},
)
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
_apply_estimated_usage(
response_json, request_body, model_obj, amount, unit, "chat"
)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="chat",
model=response_json.get("model"),
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data and "usage" in response_json:
# Inject cost breakdown into both the response body (so the
# SDK's body extractor picks up the msats breakdown) and the
# response headers (so the SDK's header extractor works too).
_inject_cost_into_usage(response_json, cost_data)
if not cost_data:
logger.error(
"Failed to calculate cost for response",
extra={
"amount": amount,
"unit": unit,
"response_model": response_json.get("model", "unknown"),
},
)
return Response(
content=json.dumps(
{
"error": {
"message": "Error forwarding request to upstream",
"type": "upstream_error",
"code": response.status_code,
}
}
),
status_code=response.status_code,
media_type="application/json",
)
response_headers = dict(response.headers)
if "transfer-encoding" in response_headers:
del response_headers["transfer-encoding"]
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
_inject_cost_response_headers(response_headers, cost_data)
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
logger.debug(
"Processing non-streaming response cost calculation",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": response_json.get("model", "unknown"),
},
)
if refund_amount > 0:
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for non-streaming response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
return Response(
content=json.dumps(response_json),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError as e:
logger.error(
"Failed to parse JSON from upstream response",
extra={
"error": str(e),
"content_preview": content_str[:200] + "..."
if len(content_str) > 200
else content_str,
"amount": amount,
"unit": unit,
},
)
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response.headers["X-Cashu"] = refund_token
await store_cashu_transaction(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=token_mint_url(refund_token, mint),
typ="out",
request_id=request_id,
)
logger.warning(
"Emergency refund issued due to JSON parse error",
extra={
"original_amount": amount,
"refund_amount": emergency_refund,
},
)
return Response(
content=content_str,
status_code=response.status_code,
headers=dict(response.headers),
media_type="application/json",
)
async def handle_x_cashu_chat_completion(
self,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse | Response:
"""Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming.
Args:
response: Response from upstream
amount: Payment amount received
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
Returns:
StreamingResponse or Response depending on response type
"""
logger.debug(
"Handling chat completion response",
extra={"amount": amount, "unit": unit, "status_code": response.status_code},
)
try:
content = await response.aread()
content_str = (
content.decode("utf-8") if isinstance(content, bytes) else content
)
is_streaming = _is_sse_body(
response.headers.get("content-type"), content_str
)
logger.debug(
"Chat completion response analysis",
extra={
"is_streaming": is_streaming,
"content_length": len(content_str),
"amount": amount,
"unit": unit,
},
)
if is_streaming:
return await self.handle_x_cashu_streaming_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
request_body=request_body,
)
else:
return await self.handle_x_cashu_non_streaming_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
request_body=request_body,
)
except Exception as e:
logger.error(
"Error processing chat completion response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"amount": amount,
"unit": unit,
},
)
return ClosingStreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=dict(response.headers),
)
async def forward_x_cashu_request(
self,
request: Request,
path: str,
headers: dict,
amount: int,
unit: str,
max_cost_for_model: int,
model_obj: Model,
mint: str | None = None,
*,
request_body: bytes | None = None,
) -> Response | StreamingResponse:
"""Forward request paid with X-Cashu token to upstream service.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
amount: Payment amount from X-Cashu token
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
model_obj: Model object for the request
Returns:
Response or StreamingResponse with refund if applicable
"""
completion_path = _openai_completion_path(path)
if path.startswith("v1/"):
path = path.replace("v1/", "")
if request_body is None:
request_body = await request.body()
if (
path.endswith("messages/count_tokens")
and not self.supports_anthropic_messages
):
result = count_tokens_locally(request_body, model_obj)
refund_token = await self.send_refund(
amount,
unit,
mint,
request_id=getattr(request.state, "request_id", None),
)
result.headers["X-Cashu"] = refund_token
return result
if (
path.endswith("messages")
and not path.endswith("count_tokens")
and not self.supports_anthropic_messages
):
return await self._forward_x_cashu_messages_via_litellm(
request_body=request_body,
amount=amount,
unit=unit,
max_cost_for_model=max_cost_for_model,
model_obj=model_obj,
mint=mint,
request_id=getattr(request.state, "request_id", None),
)
url = f"{self.base_url}/{path}"
transformed_body = self.prepare_request_body(
request_body,
model_obj,
include_stream_usage=completion_path is not None,
)
logger.debug(
"Forwarding request to upstream",
extra={
"url": url,
"method": request.method,
"path": path,
"amount": amount,
"unit": unit,
},
)
client = build_x_cashu_client()
response: httpx.Response | None = None
try:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=transformed_body if transformed_body else request_body,
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
if response.status_code != 200:
logger.error(
"Received upstream response",
extra={
"reason_phrase": response.reason_phrase,
"status_code": response.status_code,
"path": path,
"response_headers": dict(response.headers),
},
)
else:
logger.debug(
"Received upstream response",
extra={
"status_code": response.status_code,
"path": path,
"response_headers": dict(response.headers),
},
)
if response.status_code != 200:
logger.warning(
"Upstream request failed, processing refund",
extra={
"status_code": response.status_code,
"path": path,
"amount": amount,
"unit": unit,
},
)
refund_token = await self.send_refund(
amount,
unit,
mint,
request_id=getattr(request.state, "request_id", None),
)
logger.info(
"Refund processed for failed upstream 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": "Error forwarding request to upstream",
"type": "upstream_error",
# Pass the status as the code so a provider
# 4xx keeps the legacy numeric ``code``.
"code": client_code_for_upstream_error(
response.status_code, response.status_code
),
"upstream_status": response.status_code,
"refund_token": refund_token,
}
}
),
status_code=client_status_for_upstream_error(response.status_code),
media_type="application/json",
)
error_response.headers["X-Cashu"] = refund_token
error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM
await close_upstream_exchange(response, client)
return error_response
if _x_cashu_path_has_settlement_handler(path):
logger.debug(
"Processing completion/embeddings/messages response",
extra={"path": path, "amount": amount, "unit": unit},
)
result = await self.handle_x_cashu_chat_completion(
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
request_body=request_body,
)
if isinstance(result, StreamingResponse) and not response.is_closed:
return attach_upstream_stream_owner(result, response, client)
await close_upstream_exchange(response, client)
return result
logger.debug(
"Streaming non-chat response",
extra={"path": path, "status_code": response.status_code},
)
return ClosingStreamingResponse(
OwnedUpstreamStream(response.aiter_bytes(), response, client),
status_code=response.status_code,
headers=dict(response.headers),
)
except asyncio.CancelledError:
await close_upstream_exchange(response, client)
raise
except Exception as exc:
await close_upstream_exchange(response, client)
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,
},
)
return create_error_response(
"internal_error",
"An unexpected server error occurred",
500,
request=request,
)
async def handle_x_cashu_responses(
self,
request: Request,
x_cashu_token: str,
path: str,
max_cost_for_model: int,
model_obj: Model,
*,
request_body: bytes | None = None,
) -> Response | StreamingResponse:
"""Handle X-Cashu payment for Responses API requests.
Args:
request: Original FastAPI request
x_cashu_token: X-Cashu token from request header
path: Request path
max_cost_for_model: Maximum cost for the model
model_obj: Model object for the request
Returns:
Response or StreamingResponse from upstream with refund if applicable
"""
logger.debug(
"Processing X-Cashu payment for Responses API",
extra={
"path": path,
"method": request.method,
"token_preview": x_cashu_token[:20] + "..."
if len(x_cashu_token) > 20
else x_cashu_token,
},
)
redeemed = False
try:
headers = dict(request.headers)
amount, unit, mint = await recieve_token(x_cashu_token)
# Reject a zero/negative redemption (empty/dust token, or a value
# fully consumed by fees) before marking the token redeemed, so it
# classifies as cashu_token_zero_value like the bearer/top-up paths
# rather than being forwarded as a free request.
if amount <= 0:
raise ValueError(
f"Redeemed token amount must be positive, got {amount} {unit}"
)
redeemed = True
headers = self.prepare_headers(dict(request.headers))
request_id = getattr(request.state, "request_id", None)
await store_cashu_transaction(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
logger.info(
"X-Cashu token redeemed for Responses API",
extra={"amount": amount, "unit": unit, "path": path, "mint": mint},
)
return await self.forward_x_cashu_responses_request(
request,
path,
headers,
amount,
unit,
max_cost_for_model,
model_obj,
mint,
request_body=request_body,
)
except Exception as e:
error_message = str(e)
logger.error(
"X-Cashu payment for Responses API failed",
extra={
"error": error_message,
"error_type": type(e).__name__,
"path": path,
"method": request.method,
},
)
# Post-redemption the token is spent; a forwarding failure must not
# be reported as a retryable redemption error (see handle_x_cashu).
if redeemed:
upstream_status = getattr(e, "status_code", None)
upstream_code = getattr(e, "code", None)
return create_error_response(
"upstream_error",
"Payment succeeded but the upstream request failed",
client_status_for_upstream_error(upstream_status, upstream_code),
request=request,
code=client_code_for_upstream_error(upstream_status, upstream_code),
details=upstream_status_details(None, upstream_status),
error_scope=ERROR_SCOPE_UPSTREAM,
)
classified = classify_redemption_error(e)
if classified is None:
return create_error_response(
"api_error",
"Internal error during token redemption",
500,
request=request,
code="internal_error",
)
error_type, status_code, message, error_code = classified
# Echo the token back only when it is still spendable, so clients
# can recover it; a spent/consumed token is never re-offered.
echo_token = None if error_code in SPENT_TOKEN_CODES else x_cashu_token
return create_error_response(
error_type,
message,
status_code,
request=request,
token=echo_token,
code=error_code,
)
async def forward_x_cashu_responses_request(
self,
request: Request,
path: str,
headers: dict,
amount: int,
unit: str,
max_cost_for_model: int,
model_obj: Model,
mint: str | None = None,
*,
request_body: bytes | None = None,
) -> Response | StreamingResponse:
"""Forward Responses API request paid with X-Cashu token to upstream service.
Args:
request: Original FastAPI request
path: Request path
headers: Prepared headers for upstream
amount: Payment amount from X-Cashu token
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
model_obj: Model object for the request
mint: Mint URL for refund tokens
Returns:
Response or StreamingResponse with refund if applicable
"""
if path.startswith("v1/"):
path = path.replace("v1/", "")
url = f"{self.base_url}/{path}"
if request_body is None:
request_body = await request.body()
transformed_body = self.prepare_responses_request_body(request_body, model_obj)
logger.debug(
"Forwarding Responses API request to upstream with X-Cashu payment",
extra={
"url": url,
"method": request.method,
"path": path,
"amount": amount,
"unit": unit,
},
)
client = build_x_cashu_client()
response: httpx.Response | None = None
try:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=transformed_body if transformed_body else request_body,
params=self.prepare_params(path, request.query_params),
),
stream=True,
)
logger.debug(
"Received upstream Responses API response",
extra={
"status_code": response.status_code,
"path": path,
"response_headers": dict(response.headers),
},
)
if response.status_code != 200:
logger.warning(
"Upstream Responses API request failed, processing refund",
extra={
"status_code": response.status_code,
"path": path,
"amount": amount,
"unit": unit,
},
)
refund_token = await self.send_refund(
amount,
unit,
mint,
request_id=getattr(request.state, "request_id", None),
)
logger.info(
"Refund processed for failed upstream Responses API 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": "Error forwarding Responses API request to upstream",
"type": "upstream_error",
# Pass the status as the code so a provider
# 4xx keeps the legacy numeric ``code``.
"code": client_code_for_upstream_error(
response.status_code, response.status_code
),
"upstream_status": response.status_code,
"refund_token": refund_token,
}
}
),
status_code=client_status_for_upstream_error(response.status_code),
media_type="application/json",
)
error_response.headers["X-Cashu"] = refund_token
error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM
await close_upstream_exchange(response, client)
return error_response
if path.startswith("responses"):
logger.debug(
"Processing Responses API response",
extra={"path": path, "amount": amount, "unit": unit},
)
result = await self.handle_x_cashu_responses_completion(
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
request_body=request_body,
)
if isinstance(result, StreamingResponse) and not response.is_closed:
return attach_upstream_stream_owner(result, response, client)
await close_upstream_exchange(response, client)
return result
logger.debug(
"Streaming non-responses response",
extra={"path": path, "status_code": response.status_code},
)
return ClosingStreamingResponse(
OwnedUpstreamStream(response.aiter_bytes(), response, client),
status_code=response.status_code,
headers=dict(response.headers),
)
except asyncio.CancelledError:
await close_upstream_exchange(response, client)
raise
except Exception as exc:
await close_upstream_exchange(response, client)
tb = traceback.format_exc()
logger.error(
"Unexpected error in upstream Responses API 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,
},
)
return create_error_response(
"internal_error",
"An unexpected server error occurred",
500,
request=request,
)
async def handle_x_cashu_responses_completion(
self,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse | Response:
"""Handle Responses API completion response for X-Cashu payment.
Args:
response: Response from upstream
amount: Payment amount received
unit: Payment unit (sat or msat)
max_cost_for_model: Maximum cost for the model
mint: Mint URL for refund tokens
Returns:
StreamingResponse or Response depending on response type
"""
logger.debug(
"Handling Responses API completion response",
extra={"amount": amount, "unit": unit, "status_code": response.status_code},
)
try:
content = await response.aread()
content_str = (
content.decode("utf-8") if isinstance(content, bytes) else content
)
is_streaming = _is_sse_body(
response.headers.get("content-type"), content_str
)
logger.debug(
"Responses API completion response analysis",
extra={
"is_streaming": is_streaming,
"content_length": len(content_str),
"amount": amount,
"unit": unit,
},
)
if is_streaming:
return await self.handle_x_cashu_streaming_responses_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
request_body=request_body,
)
else:
return await self.handle_x_cashu_non_streaming_responses_response(
content_str,
response,
amount,
unit,
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
request_body=request_body,
)
except Exception as e:
logger.error(
"Error processing Responses API completion response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"amount": amount,
"unit": unit,
},
)
return ClosingStreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=dict(response.headers),
)
async def handle_x_cashu_streaming_responses_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API response for X-Cashu payment.
Similar to regular streaming but handles Responses API specific tokens like reasoning_tokens.
"""
events = _parse_sse_events(content_str)
logger.debug(
"Processing streaming Responses API response",
extra={
"amount": amount,
"unit": unit,
"event_count": len(events),
},
)
response_headers = dict(response.headers)
if "transfer-encoding" in response_headers:
del response_headers["transfer-encoding"]
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
usage_data: dict | None = None
model: str | None = None
reasoning_tokens = 0
cost_data: CostData | MaxCostData | None = None
usage_estimator = MissingUsageEstimator(request_body, model_obj)
for _fields, data in events:
if data.strip() == "[DONE]":
continue
try:
data_json = json.loads(data)
except json.JSONDecodeError:
continue
if not isinstance(data_json, dict):
continue
usage_estimator.observe(data_json)
# Canonical Responses API events carry model and usage nested under
# "response" (response.completed/incomplete); older shapes put them
# at the top level.
payload = _responses_usage_payload(data_json)
if isinstance(payload.get("usage"), dict):
usage_data = payload["usage"]
model = payload.get("model") or model
details = usage_data.get("output_tokens_details")
if isinstance(details, dict):
reasoning_tokens = details.get("reasoning_tokens", 0)
elif "reasoning_tokens" in usage_data:
reasoning_tokens = usage_data["reasoning_tokens"]
elif not model and payload.get("model"):
model = payload["model"]
if not usage_data:
usage_data = usage_estimator.estimated_usage(model)
if usage_data:
logger.warning(
"No usage in streaming Responses API response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"estimated_usage": usage_data,
},
)
else:
logger.debug(
"Found usage data in streaming Responses API response",
extra={
"model": model,
"usage_data": usage_data,
"reasoning_tokens": reasoning_tokens,
"amount": amount,
"unit": unit,
},
)
response_data = {"usage": usage_data, "model": model or "unknown"}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="responses.streaming",
model=model,
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
if refund_amount > 0:
logger.debug(
"Processing refund for streaming Responses API response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
"reasoning_tokens": reasoning_tokens,
},
)
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming Responses API response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
else:
logger.debug(
"No refund needed for streaming Responses API response",
extra={
"amount": amount,
"cost_msats": cost_data.total_msats,
"model": model,
},
)
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e:
logger.error(
"Error calculating cost for streaming Responses API response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"amount": amount,
"unit": unit,
},
)
provider_seen: str | None = None
for i, (fields, data) in enumerate(events):
if data.strip() == "[DONE]":
continue
try:
data_json = json.loads(data)
except json.JSONDecodeError:
continue
if not isinstance(data_json, dict):
continue
provider_before = data_json.get("provider")
provider_seen = self._stamp_streamed_provider(data_json, provider_seen)
changed = data_json.get("provider") != provider_before
payload = _responses_usage_payload(data_json)
if cost_data and isinstance(payload.get("usage"), dict):
_inject_cost_into_usage(payload, cost_data)
changed = True
if changed:
events[i] = (fields, json.dumps(data_json))
async def generate() -> AsyncGenerator[bytes, None]:
for fields, data in events:
yield _render_sse_event(fields, data).encode("utf-8")
return ClosingStreamingResponse(
generate(),
status_code=response.status_code,
headers=response_headers,
media_type="text/plain",
)
async def handle_x_cashu_non_streaming_responses_response(
self,
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> Response:
"""Handle non-streaming Responses API response for X-Cashu payment."""
logger.debug(
"Processing non-streaming Responses API response",
extra={"amount": amount, "unit": unit, "content_length": len(content_str)},
)
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
_apply_estimated_usage(
response_json, request_body, model_obj, amount, unit, "responses"
)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="responses",
model=response_json.get("model"),
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data and "usage" in response_json:
_inject_cost_into_usage(response_json, cost_data)
if not cost_data:
logger.error(
"Failed to calculate cost for Responses API response",
extra={
"amount": amount,
"unit": unit,
"response_model": response_json.get("model", "unknown"),
},
)
return Response(
content=json.dumps(
{
"error": {
"message": "Error forwarding Responses API request to upstream",
"type": "upstream_error",
"code": response.status_code,
}
}
),
status_code=response.status_code,
media_type="application/json",
)
response_headers = dict(response.headers)
if "transfer-encoding" in response_headers:
del response_headers["transfer-encoding"]
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
_inject_cost_response_headers(response_headers, cost_data)
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
logger.debug(
"Processing non-streaming Responses API cost calculation",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": response_json.get("model", "unknown"),
},
)
if refund_amount > 0:
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for non-streaming Responses API response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
return Response(
content=json.dumps(response_json),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError as e:
logger.error(
"Failed to parse JSON from upstream Responses API response",
extra={
"error": str(e),
"content_preview": content_str[:200] + "..."
if len(content_str) > 200
else content_str,
"amount": amount,
"unit": unit,
},
)
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response.headers["X-Cashu"] = refund_token
await store_cashu_transaction(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=token_mint_url(refund_token, mint),
typ="out",
request_id=request_id,
)
logger.warning(
"Emergency refund issued for Responses API due to JSON parse error",
extra={
"original_amount": amount,
"refund_amount": emergency_refund,
},
)
return Response(
content=content_str,
status_code=response.status_code,
headers=dict(response.headers),
media_type="application/json",
)
async def handle_x_cashu(
self,
request: Request,
x_cashu_token: str,
path: str,
max_cost_for_model: int,
model_obj: Model,
*,
request_body: bytes | None = None,
) -> Response | StreamingResponse:
"""Handle request with X-Cashu token payment, redeeming token and forwarding request.
Args:
request: Original FastAPI request
x_cashu_token: X-Cashu token from request header
path: Request path
max_cost_for_model: Maximum cost for the model
model_obj: Model object for the request
Returns:
Response or StreamingResponse from upstream with refund if applicable
"""
logger.debug(
"Processing X-Cashu payment request",
extra={
"path": path,
"method": request.method,
"token_preview": x_cashu_token[:20] + "..."
if len(x_cashu_token) > 20
else x_cashu_token,
},
)
# Reject before redemption so the client keeps its token.
if not _x_cashu_path_has_settlement_handler(path):
logger.warning(
"Rejecting X-Cashu request for unsupported endpoint",
extra={"path": path, "method": request.method},
)
return create_error_response(
"invalid_request_error",
"X-Cashu payment is not supported on this endpoint; use bearer "
"(deposit) authentication instead. The token was not redeemed.",
400,
request=request,
code="x_cashu_unsupported_endpoint",
)
redeemed = False
try:
headers = dict(request.headers)
amount, unit, mint = await recieve_token(x_cashu_token)
# Reject a zero/negative redemption (empty/dust token, or a value
# fully consumed by fees) before marking the token redeemed, so it
# classifies as cashu_token_zero_value like the bearer/top-up paths
# rather than being forwarded as a free request.
if amount <= 0:
raise ValueError(
f"Redeemed token amount must be positive, got {amount} {unit}"
)
redeemed = True
headers = self.prepare_headers(dict(request.headers))
request_id = getattr(request.state, "request_id", None)
await store_cashu_transaction(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
logger.info(
"X-Cashu token redeemed successfully",
extra={"amount": amount, "unit": unit, "path": path, "mint": mint},
)
return await self.forward_x_cashu_request(
request,
path,
headers,
amount,
unit,
max_cost_for_model,
model_obj,
mint,
request_body=request_body,
)
except Exception as e:
error_message = str(e)
logger.error(
"X-Cashu payment request failed",
extra={
"error": error_message,
"error_type": type(e).__name__,
"path": path,
"method": request.method,
},
)
# Once redeemed the token is spent, so a later forwarding failure
# must not surface as a retryable mint_unreachable (spent-token retry
# bait). Redemption classification only applies while not redeemed.
if redeemed:
upstream_status = getattr(e, "status_code", None)
upstream_code = getattr(e, "code", None)
return create_error_response(
"upstream_error",
"Payment succeeded but the upstream request failed",
client_status_for_upstream_error(upstream_status, upstream_code),
request=request,
code=client_code_for_upstream_error(upstream_status, upstream_code),
details=upstream_status_details(None, upstream_status),
error_scope=ERROR_SCOPE_UPSTREAM,
)
classified = classify_redemption_error(e)
if classified is None:
return create_error_response(
"api_error",
"Internal error during token redemption",
500,
request=request,
code="internal_error",
)
error_type, status_code, message, error_code = classified
# Echo the token back only when it is still spendable, so clients
# can recover it; a spent/consumed token is never re-offered.
echo_token = None if error_code in SPENT_TOKEN_CODES else x_cashu_token
return create_error_response(
error_type,
message,
status_code,
request=request,
token=echo_token,
code=error_code,
)
def _apply_provider_fee_to_model(self, model: Model) -> Model:
"""Apply provider fee to model's USD pricing and calculate max costs.
Cache rates missing from the upstream pricing feed are backfilled from
litellm's cost map first, so they carry the provider fee like every
other price component.
Args:
model: Model object to update
Returns:
Model with provider fee applied to pricing and max costs calculated
"""
base_pricing = backfill_cache_pricing(model.id, model.pricing)
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in base_pricing.dict().items()}
)
temp_model = model.copy(
update={"pricing": adjusted_pricing, "sats_pricing": None}
)
(
adjusted_pricing.max_prompt_cost,
adjusted_pricing.max_completion_cost,
adjusted_pricing.max_cost,
) = _calculate_usd_max_costs(temp_model)
return model.copy(update={"pricing": adjusted_pricing})
async def fetch_models(self) -> list[Model]:
"""Fetch available models from upstream API and update cache.
Returns:
List of Model objects with pricing
"""
try:
or_models, provider_models_response = await asyncio.gather(
self._fetch_openrouter_models(),
self._fetch_provider_models(),
)
provider_model_ids = self._parse_model_ids(provider_models_response)
found_models = []
not_found_models = []
for model_id in provider_model_ids:
or_model = self._match_model(model_id, or_models)
if or_model:
try:
model = Model(**or_model) # type: ignore
found_models.append(model)
except Exception as e:
logger.warning(
f"Failed to parse model {model_id}",
extra={"error": str(e), "error_type": type(e).__name__},
)
else:
not_found_models.append(model_id)
if not_found_models:
logger.debug(
f"({len(not_found_models)}/{len(provider_model_ids)}) unmatched models for {self.provider_type or self.base_url}",
extra={"not_found_models": not_found_models},
)
return found_models
except Exception as e:
logger.error(
f"Error fetching models for {self.provider_type or self.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
)
return []
async def _fetch_openrouter_models(self) -> list[dict]:
"""Fetch models from OpenRouter API."""
url = "https://openrouter.ai/api/v1/models"
embeddings_url = "https://openrouter.ai/api/v1/embeddings/models"
async with httpx.AsyncClient(timeout=30.0) as client:
models_response, embeddings_response = await asyncio.gather(
client.get(url), client.get(embeddings_url), return_exceptions=True
)
all_models = []
def process_models_response(
response: httpx.Response | BaseException,
) -> list[dict]:
if not isinstance(response, BaseException):
response.raise_for_status()
data = response.json()
return [
model
for model in data.get("data", [])
if ":free" not in model.get("id", "").lower()
]
return []
all_models.extend(process_models_response(models_response))
all_models.extend(process_models_response(embeddings_response))
return all_models
async def _fetch_provider_models(self) -> dict:
"""Fetch models from provider's API."""
url = f"{self.base_url.rstrip('/')}/models"
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else None
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url, headers=headers)
response.raise_for_status()
return response.json()
def _parse_model_ids(self, response: dict) -> list[str]:
"""Parse model IDs from provider response."""
return [model.get("id") for model in response.get("data", []) if "id" in model]
def _match_model(self, model_id: str, or_models: list[dict]) -> dict | None:
"""Match provider model ID with OpenRouter model."""
return next(
(
model
for model in or_models
if (model.get("id") == model_id)
or (model.get("id", "").split("/")[-1] == model_id)
or (model.get("canonical_slug") == model_id)
or (model.get("canonical_slug", "").split("/")[-1] == model_id)
),
None,
)
async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API."""
try:
async with create_session() as session:
provider = (
await session.get(UpstreamProviderRow, self.db_id)
if self.db_id is not None
else None
)
if not provider or not provider.id:
raise HTTPException(status_code=404, detail="Provider not found")
db_models = await list_models(
session=session,
upstream_id=provider.id,
include_disabled=False,
apply_fees=False,
)
db_model_ids: set[str] = {model.id for model in db_models}
models = await self.fetch_models()
model_ids = [model.id for model in models]
diff = set(db_model_ids) - set(model_ids)
for db_model_id in diff:
found_db_model = next(
(
model_obj
for model_obj in db_models
if model_obj.id == db_model_id
)
)
models.append(found_db_model)
models_with_fees = [
self._apply_provider_fee_to_model(m) for m in models
]
try:
sats_to_usd = sats_usd_price()
self._models_cache = [
_update_model_sats_pricing(m, sats_to_usd)
for m in models_with_fees
]
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {
m.forwarded_model_id or m.id: m for m in self._models_cache
}
except Exception as e:
logger.error(
f"Failed to refresh models cache for {self.provider_type or self.base_url}",
extra={"error": repr(e), "error_type": type(e).__name__},
)
def get_cached_models(self) -> list[Model]:
"""Get cached models for this provider.
Returns:
List of cached Model objects
"""
return self._models_cache
def get_cached_model_by_id(self, model_id: str) -> Model | None:
"""Get a specific cached model by ID.
Args:
model_id: Model identifier
Returns:
Model object or None if not found
"""
return self._models_by_id.get(model_id)
@classmethod
async def create_account_static(cls) -> dict[str, object]:
"""Create a new account with the provider (class method, no instance needed).
Returns:
Dict with account creation details including api_key
Raises:
NotImplementedError: If provider does not support account creation
"""
raise NotImplementedError(
f"Provider {cls.provider_type} does not support account creation"
)
async def create_account(self) -> dict[str, object]:
"""Create a new account with the provider.
Returns:
Dict with account creation details including api_key
Raises:
NotImplementedError: If provider does not support account creation
"""
raise NotImplementedError(
f"Provider {self.provider_type} does not support account creation"
)
async def initiate_topup(self, amount: int) -> TopupData:
"""Initiate a Lightning Network top-up for the provider account.
Args:
amount: Amount in currency units to top up
Returns:
TopupData with standardized invoice information
Raises:
NotImplementedError: If provider does not support top-up
"""
raise NotImplementedError(
f"Provider {self.provider_type} does not support top-up"
)
async def get_balance(self) -> float | None:
"""Get the current account balance from the provider.
Returns:
Float representing the balance amount, or None if not supported/available.
Typically in USD or the provider's credit unit.
Raises:
NotImplementedError: If provider does not support balance checking (default behavior)
"""
raise NotImplementedError(
f"Provider {self.provider_type} does not support balance checking"
)