mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
731 lines
25 KiB
Python
731 lines
25 KiB
Python
import asyncio
|
|
import base64
|
|
import ipaddress
|
|
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 ..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)
|
|
|
|
|
|
def _is_blocked_address(address: str) -> bool:
|
|
"""Allow only globally reachable addresses (RFC 6890)."""
|
|
try:
|
|
ip = ipaddress.ip_address(address)
|
|
except ValueError:
|
|
return True
|
|
if isinstance(ip, ipaddress.IPv6Address):
|
|
# An embedded v4 address would otherwise smuggle a rejected target past
|
|
# the v6 checks.
|
|
for embedded in (ip.ipv4_mapped, ip.sixtofour):
|
|
if embedded is not None:
|
|
return _is_blocked_address(str(embedded))
|
|
return not ip.is_global or ip.is_multicast
|
|
|
|
|
|
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,
|
|
)
|