Files
routstr-core/routstr/payment/helpers.py

716 lines
24 KiB
Python

import asyncio
import base64
import json
import math
import socket
from io import BytesIO
from typing import Any
from urllib.parse import urlsplit, urlunsplit
import httpx
from fastapi import HTTPException, Response
from fastapi.requests import Request
from PIL import Image
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger
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 ..core.settings import settings
from ..net_guard import is_blocked_address as _is_blocked_address
from ..wallet import (
UntrustedSourceMintError,
classify_redemption_error,
deserialize_token_from_string,
is_trusted_source_mint,
)
from .responses_input import (
FILE_ID_URL_PREFIX,
count_input_images,
input_image_part_to_image_url,
responses_input_to_messages,
)
logger = get_logger(__name__)
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
if x_cashu := headers.get("x-cashu", None):
cashu_token = x_cashu
logger.debug(
"Using X-Cashu token",
extra={
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token
},
)
elif auth := headers.get("authorization", None):
logger.debug(
"Skipping preflight token balance check for Authorization header",
extra={
"auth_preview": auth[:20] + "..." if len(auth) > 20 else auth,
},
)
return
else:
logger.error("No authentication token provided")
raise HTTPException(status_code=401, detail="Unauthorized")
# Handle empty token
if not cashu_token:
logger.error("Empty token provided")
raise HTTPException(
status_code=401,
detail={
"error": {
"message": "API key or Cashu token required",
"type": "invalid_request_error",
"code": "missing_api_key",
}
},
)
# Handle regular API keys (sk-*)
if cashu_token.startswith("sk-"):
return
try:
token_obj = deserialize_token_from_string(cashu_token)
except Exception:
# Invalid token format - let the auth system handle it
raise HTTPException(
status_code=401,
detail="Invalid authentication token format",
)
if not is_trusted_source_mint(token_obj.mint):
classified = classify_redemption_error(
UntrustedSourceMintError(f"Untrusted source mint: {token_obj.mint}")
)
assert classified is not None
error_type, status_code, message, error_code = classified
raise HTTPException(
status_code=status_code,
detail={
"error": {"message": message, "type": error_type, "code": error_code}
},
)
amount_msat = (
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
)
if max_cost_for_model > amount_msat:
raise HTTPException(
status_code=402,
detail={
"reason": "Insufficient balance",
"amount_required_msat": max_cost_for_model,
"model": body.get("model", "unknown"),
"type": "minimum_balance_required",
},
)
async def get_max_cost_for_model(
model: str,
session: AsyncSession,
model_obj: Any | None = None,
) -> int:
"""Get the maximum cost for a specific model from providers with overrides."""
logger.debug(
"Getting max cost for model",
extra={
"model": model,
"fixed_pricing": settings.fixed_pricing,
},
)
if settings.fixed_pricing:
default_cost_msats = settings.fixed_cost_per_request * 1000
logger.debug(
"Using fixed cost pricing",
extra={"cost_msats": default_cost_msats, "model": model},
)
return max(settings.min_request_msat, default_cost_msats)
if not model_obj:
from ..proxy import get_model_instance
model_obj = get_model_instance(model)
if not model_obj:
fallback_msats = settings.fixed_cost_per_request * 1000
logger.warning(
"Model not found in providers or overrides",
extra={
"requested_model": model,
"using_default_cost": fallback_msats,
},
)
return max(settings.min_request_msat, fallback_msats)
if model_obj.sats_pricing:
try:
max_cost = (
model_obj.sats_pricing.max_cost
* 1000
* (1 - settings.tolerance_percentage / 100)
)
logger.debug(
"Found model-specific max cost",
extra={"model": model, "max_cost_msats": max_cost},
)
calculated_msats = int(max_cost)
return max(settings.min_request_msat, calculated_msats)
except Exception as e:
logger.error(
"Error calculating max cost from model pricing",
extra={"model": model, "error": str(e)},
)
logger.warning(
"Model pricing not found, using fixed cost",
extra={
"model": model,
"default_cost_msats": settings.fixed_cost_per_request * 1000,
},
)
return max(settings.min_request_msat, settings.fixed_cost_per_request * 1000)
async def calculate_discounted_max_cost(
max_cost_for_model: int,
body: dict,
model_obj: Any | None = None,
) -> int:
"""Calculate the discounted max cost for a request using model pricing when available.
Completion discounts are trimmed from the largest declared cap among
``max_tokens`` and ``max_completion_tokens`` (chat/completions) or
``max_output_tokens`` (responses).
"""
if settings.fixed_pricing:
return max_cost_for_model
model = body.get("model", "unknown")
model_pricing = model_obj.sats_pricing if model_obj else None
if not model_pricing:
return max_cost_for_model
tol = settings.tolerance_percentage
tol_factor = max(0.0, 1 - float(tol) / 100.0)
max_prompt_allowed_sats = model_pricing.max_prompt_cost * tol_factor
max_completion_allowed_sats = model_pricing.max_completion_cost * tol_factor
if model_obj:
prompt_token_limit: int | None = None
if model_obj.top_provider and (
model_obj.top_provider.context_length
or model_obj.top_provider.max_completion_tokens
):
cl = model_obj.top_provider.context_length
mct = model_obj.top_provider.max_completion_tokens
if cl and mct:
prompt_token_limit = max(0, cl - mct)
elif cl:
prompt_token_limit = cl
elif mct:
prompt_token_limit = 0
elif model_obj.context_length:
prompt_token_limit = model_obj.context_length
if prompt_token_limit is not None:
max_prompt_allowed_sats = (
prompt_token_limit * model_pricing.prompt * tol_factor
)
adjusted = max_cost_for_model
messages = body.get("messages")
# Estimated over the whole body: a discount driven by message text alone lets
# a caller hide prompt weight elsewhere, shrink the reservation, and be billed
# for work the reservation never covered.
prompt_tokens = estimate_prompt_tokens(body)
# Images are billed as tokens by the upstream but carry no text for
# ``estimate_prompt_tokens`` to count, so they are estimated separately and
# added on both the chat (``messages``) and Responses (``input``) paths.
image_tokens = 0
if isinstance(messages, list):
image_tokens += await estimate_image_tokens_in_messages(messages)
input_data = body.get("input")
if isinstance(input_data, list):
converted = responses_input_to_messages(input_data)
if converted is None:
image_tokens += count_input_images(input_data) * _MAX_ORIGINAL_IMAGE_TOKENS
else:
image_tokens += await estimate_image_tokens_in_messages(converted)
if image_tokens > 0:
logger.debug(
"Found images in request",
extra={
"model": model,
"image_tokens": image_tokens,
},
)
prompt_tokens += image_tokens
if prompt_tokens > 0:
estimated_prompt_delta_sats = (
max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt
)
if estimated_prompt_delta_sats > 0:
adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000)
# Completion caps arrive under several names: ``max_tokens`` (legacy
# chat), ``max_completion_tokens`` (modern chat) and ``max_output_tokens``
# (Responses API). When a request declares more than one, reserve against
# the largest: upstream precedence between the fields varies by provider,
# so the smaller cap may not be honored and the reservation must never
# under-cover what the upstream could bill.
max_tokens_int: int | None = None
for cap_field in ("max_tokens", "max_completion_tokens", "max_output_tokens"):
cap_raw = body.get(cap_field)
if cap_raw is None:
continue
try:
cap_int = int(cap_raw)
except (TypeError, ValueError):
logger.warning(
"Invalid completion token cap; ignoring in cost adjustment",
extra={
"field": cap_field,
"value": str(cap_raw)[:64],
"model": model,
},
)
continue
max_tokens_int = (
cap_int if max_tokens_int is None else max(max_tokens_int, cap_int)
)
if max_tokens_int is not None:
estimated_completion_delta_sats = (
max_completion_allowed_sats - max_tokens_int * model_pricing.completion
)
if estimated_completion_delta_sats > 0:
adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000)
logger.debug(
"Discounted max cost computed",
extra={
"model": model,
"original_msats": max_cost_for_model,
"adjusted_msats": adjusted,
"tolerance_pct": tol,
},
)
return max(settings.min_request_msat, adjusted)
def estimate_tokens(messages: list) -> int:
"""Estimate tokens for text content, excluding image_url fields."""
total = 0
for msg in messages:
if isinstance(msg, dict):
content = msg.get("content")
if isinstance(content, str):
total += len(content)
elif isinstance(content, list):
total += sum(
len(item.get("text", ""))
for item in content
if isinstance(item, dict) and item.get("type") == "text"
)
return total // 3
def _sum_string_chars(node: Any) -> int:
"""Recursively sum the length of every string in the tree, keys included.
Nothing is excluded. Keys count because JSON-schema property names are
forwarded to the provider, and no exclusion rule can be trusted here: every
part of the body is caller-controlled, so any carve-out (by key name or by
value shape) is a place to hide prompt weight for free. Inline image data is
therefore counted as text too, which only makes the discount smaller.
"""
if isinstance(node, str):
return len(node)
if isinstance(node, dict):
return sum(
len(str(key)) + _sum_string_chars(value) for key, value in node.items()
)
if isinstance(node, list):
return sum(_sum_string_chars(item) for item in node)
return 0
def _count_prompt_token_ids(node: Any) -> int:
if isinstance(node, int) and not isinstance(node, bool):
return 1
if isinstance(node, list):
return sum(_count_prompt_token_ids(item) for item in node)
return 0
def estimate_prompt_tokens(body: dict) -> int:
"""Conservatively estimate prompt tokens for the whole provider-bound body.
Every string counts, as do token IDs in legacy ``prompt`` arrays, so no
forwarded field can hide prompt weight and shrink its reservation.
"""
return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt"))
IMAGE_FETCH_TIMEOUT_SECONDS = 10.0
# Dimensions live in the header, so a prefix suffices and an endless body cannot
# pin memory.
IMAGE_FETCH_MAX_BYTES = 512 * 1024
# Fetches are sequential, so an unbounded URL list is a request-time amplifier.
IMAGE_FETCH_MAX_PER_REQUEST = 8
def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
"""Extract image dimensions from image bytes."""
try:
img = Image.open(BytesIO(image_data))
return img.size
except Exception as e:
logger.warning(
"Failed to get image dimensions, using default",
extra={"error": str(e)},
)
return (512, 512)
async def _validated_fetch_target(url: str) -> tuple[str, str]:
"""Return the URL to request and its ``Host`` header.
Cost estimation runs on the unauthenticated request body, so a caller can
otherwise aim the node at internal hosts. HTTP is rewritten to the resolved
address so the name cannot rebind between check and connect; HTTPS keeps its
hostname because certificate validation already binds the connection.
"""
parts = urlsplit(url)
if parts.scheme not in ("http", "https"):
raise ValueError(f"unsupported scheme: {parts.scheme or 'none'}")
host = parts.hostname
if not host:
raise ValueError("missing host")
default_port = 443 if parts.scheme == "https" else 80
port = parts.port or default_port
host_header = f"[{host}]" if ":" in host else host
if parts.port is not None:
host_header = f"{host_header}:{parts.port}"
infos = await asyncio.get_running_loop().getaddrinfo(
host, port, proto=socket.IPPROTO_TCP
)
if not infos:
raise ValueError("host did not resolve")
for info in infos:
if _is_blocked_address(str(info[4][0])):
raise ValueError("host resolves to a blocked address")
if parts.scheme == "https":
return url, host_header
family, _, _, _, sockaddr = infos[0]
address = str(sockaddr[0])
pinned = f"[{address}]" if family == socket.AF_INET6 else address
if parts.port is not None:
pinned = f"{pinned}:{parts.port}"
return urlunsplit((parts.scheme, pinned, parts.path, parts.query, "")), host_header
async def _fetch_image_from_url(url: str) -> bytes | None:
"""Fetch the leading bytes of an image, enough to read its dimensions."""
try:
target, host_header = await _validated_fetch_target(url)
async with httpx.AsyncClient(
timeout=IMAGE_FETCH_TIMEOUT_SECONDS, follow_redirects=False
) as client:
async with client.stream(
"GET", target, headers={"Host": host_header}
) as response:
response.raise_for_status()
chunks: list[bytes] = []
downloaded = 0
async for chunk in response.aiter_bytes():
chunks.append(chunk)
downloaded += len(chunk)
if downloaded >= IMAGE_FETCH_MAX_BYTES:
break
return b"".join(chunks)[:IMAGE_FETCH_MAX_BYTES]
except Exception as e:
logger.warning(
"Failed to fetch image from URL",
extra={"error": str(e), "url": url[:100]},
)
return None
# Patch-based image pricing (OpenAI ``detail: "original"``): the image is
# covered with 32x32px patches and billed as ceil(patches * multiplier)
# tokens, with no 512px-tile downscaling. The API rejects images above
# 30,000 patches, so at the 1.2x multiplier documented for the
# original-capable model families (gpt-5.4/5.5/5.6) the worst case a
# single image can bill is 36,000 tokens.
_IMAGE_PATCH_PX = 32
_MAX_IMAGE_PATCHES = 30_000
_MAX_ORIGINAL_IMAGE_TOKENS = (_MAX_IMAGE_PATCHES * 6 + 4) // 5 # 36,000
def _calculate_original_image_tokens(width: int, height: int) -> int:
"""Estimate tokens for an image billed at ``detail: "original"``.
Patch-based models cover the image with 32x32px patches and bill
``ceil(patches * 1.2)`` tokens. The estimate is bounded by the
30,000-patch rejection limit, which is more conservative than the
per-model resizing patch budgets (e.g. 10,000 patches on gpt-5.4/5.5)
so it never under-reserves.
"""
patches = ((width + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX) * (
(height + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX
)
bounded = min(patches, _MAX_IMAGE_PATCHES)
return (bounded * 6 + 4) // 5 # ceil(bounded * 1.2) in exact integer math
def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int:
"""Calculate image tokens based on OpenAI's vision pricing.
For low detail: 85 tokens
For high detail/auto: 85 base tokens + 170 tokens per 512px tile
For original detail: patch-based pricing at the original resolution
"""
if detail == "low":
return 85
if detail == "original":
return _calculate_original_image_tokens(width, height)
if width > 2048 or height > 2048:
aspect_ratio = width / height
if width > height:
width = 2048
height = int(width / aspect_ratio)
else:
height = 2048
width = int(height * aspect_ratio)
if width > 768 or height > 768:
aspect_ratio = width / height
if width > height:
width = 768
height = int(width / aspect_ratio)
else:
height = 768
width = int(height * aspect_ratio)
tiles_width = (width + 511) // 512
tiles_height = (height + 511) // 512
num_tiles = tiles_width * tiles_height
return 85 + (170 * num_tiles)
async def estimate_image_tokens_in_messages(messages: list) -> int:
"""Estimate total tokens for all images in messages.
Supports both base64 encoded images and image URLs.
"""
total_image_tokens = 0
fetches = 0
for message in messages:
if not isinstance(message, dict):
continue
content = message.get("content")
if not content:
continue
if isinstance(content, str):
continue
if not isinstance(content, list):
continue
for content_item in content:
if not isinstance(content_item, dict):
continue
content_type = content_item.get("type")
if content_type == "input_image":
content_item = input_image_part_to_image_url(content_item)
elif content_type != "image_url":
continue
image_url_data = content_item.get("image_url")
if not image_url_data:
continue
if isinstance(image_url_data, str):
url = image_url_data
detail = "auto"
elif isinstance(image_url_data, dict):
url = image_url_data.get("url", "")
detail = image_url_data.get("detail") or "auto"
else:
continue
if not url:
continue
if url.startswith("data:image/"):
total_image_tokens += _data_url_image_tokens(url, detail)
elif url.startswith(FILE_ID_URL_PREFIX):
total_image_tokens += _worst_case_image_tokens(detail)
elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST:
logger.warning(
"Skipping image URL fetch above per-request limit",
extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST},
)
total_image_tokens += _worst_case_image_tokens(detail)
else:
fetches += 1
image_bytes_or_none = await _fetch_image_from_url(url)
total_image_tokens += _image_bytes_tokens(
image_bytes_or_none, detail, source=url[:100]
)
return total_image_tokens
def _worst_case_image_tokens(detail: str) -> int:
"""Dimensions unknown: reserve the most ``detail`` can bill."""
if detail == "original":
return _MAX_ORIGINAL_IMAGE_TOKENS
return _calculate_image_tokens(2048, 2048, detail)
def _data_url_image_tokens(url: str, detail: str) -> int:
try:
_, base64_data = url.split(",", 1)
image_bytes = base64.b64decode(base64_data, validate=True)
except Exception as e:
logger.warning("Failed to decode base64 image", extra={"error": str(e)})
return _worst_case_image_tokens(detail)
return _image_bytes_tokens(image_bytes, detail, source="data-url")
def _image_bytes_tokens(image_bytes: bytes | None, detail: str, source: str) -> int:
if not image_bytes:
return _worst_case_image_tokens(detail)
try:
width, height = Image.open(BytesIO(image_bytes)).size
except Exception as e:
logger.warning(
"Failed to read image dimensions",
extra={"error": str(e), "source": source},
)
return _worst_case_image_tokens(detail)
tokens = _calculate_image_tokens(width, height, detail)
logger.debug(
"Calculated image tokens",
extra={
"source": source,
"width": width,
"height": height,
"detail": detail,
"tokens": tokens,
},
)
return tokens
def create_error_response(
error_type: str,
message: str,
status_code: int,
request: Request,
token: str | None = None,
code: str | int | None = None,
details: dict[str, object] | None = None,
error_scope: str | None = None,
) -> Response:
"""Create a standardized error response.
``code`` is a stable, machine-readable classification (e.g.
``UPSTREAM_RATE_LIMIT``); when omitted it defaults to the HTTP status code
for backwards compatibility. ``details`` carries optional structured,
redaction-safe context. ``error_scope`` is sent as the
:data:`ERROR_SCOPE_HEADER` response header.
"""
error_obj: dict[str, object] = {
"message": redact_org_ids(message),
"type": error_type,
"code": code if code is not None else status_code,
}
if details is not None:
error_obj["details"] = details
headers: dict[str, str] = {}
if token:
headers["X-Cashu"] = token
if error_scope is not None:
headers[ERROR_SCOPE_HEADER] = error_scope
return Response(
content=json.dumps(
{
"error": error_obj,
"request_id": getattr(request.state, "request_id", "unknown"),
}
),
status_code=status_code,
media_type="application/json",
headers=headers,
)
def create_upstream_error_response(
error: UpstreamError,
request: Request,
fallback_status: int = 502,
) -> Response:
"""Build an error response from an :class:`UpstreamError`.
Upstream-scoped errors are mapped via :mod:`routstr.core.error_scope`;
node-scoped errors keep their own status.
"""
status_code = error.status_code or fallback_status
code = getattr(error, "code", None)
details = getattr(error, "details", None)
if getattr(error, "scope", ERROR_SCOPE_UPSTREAM) == ERROR_SCOPE_NODE:
return create_error_response(
"upstream_error",
str(error),
status_code,
request=request,
code=code,
details=details,
)
return create_error_response(
"upstream_error",
str(error),
client_status_for_upstream_error(status_code, code),
request=request,
code=client_code_for_upstream_error(status_code, code),
details=upstream_status_details(details, status_code),
error_scope=ERROR_SCOPE_UPSTREAM,
)