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

677 lines
24 KiB
Python

import asyncio
import json
import random
import httpx
from fastapi import APIRouter, Depends, HTTPException
from pydantic.v1 import BaseModel, validator
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_price
from .rates import BILLABLE_PRICING_FIELDS, coerce_rate, is_usable_rate
logger = get_logger(__name__)
models_router = APIRouter()
class Architecture(BaseModel):
modality: str
input_modalities: list[str]
output_modalities: list[str]
tokenizer: str
instruct_type: str | None
class Pricing(BaseModel):
prompt: float
completion: float
request: float = 0.0
image: float = 0.0
web_search: float = 0.0
internal_reasoning: float = 0.0
input_cache_read: float = 0.0
input_cache_write: float = 0.0
max_prompt_cost: float = 0.0 # in sats not msats
max_completion_cost: float = 0.0 # in sats not msats
max_cost: float = 0.0 # in sats not msats
# The rates ``Pricing`` declares without a default, derived from the model so the
# two cannot drift. A payload that omits one writes a row that will not parse.
REQUIRED_PRICING_FIELDS = tuple(
name for name, field in Pricing.__fields__.items() if field.required
)
def has_usable_pricing(pricing: Pricing) -> bool:
"""True if every billable rate is a number a request could be billed on.
Free is usable — a rate of zero is a real price. This asks only whether the
price is well-formed. One unusable rate disqualifies the whole price even
alongside a valid one, since a request can bill on the bad field: a positive
``completion`` must not hide a negative ``prompt``.
"""
return all(
is_usable_rate(getattr(pricing, field)) for field in BILLABLE_PRICING_FIELDS
)
class TopProvider(BaseModel):
context_length: int | None = None
max_completion_tokens: int | None = None
is_moderated: bool | None = None
class Reasoning(BaseModel):
"""Per-model reasoning-effort metadata, matching OpenRouter's shape."""
mandatory: bool | None = None
default_enabled: bool | None = None
supported_efforts: list[str] | None = None
default_effort: str | None = None
supports_max_tokens: bool | None = None
class Config:
extra = "ignore"
def is_empty(self) -> bool:
return not any(
(
self.mandatory is not None,
self.default_enabled is not None,
self.supported_efforts,
self.default_effort,
self.supports_max_tokens is not None,
)
)
class Model(BaseModel):
id: str
name: str
created: int
description: str
context_length: int
architecture: Architecture
pricing: Pricing
sats_pricing: Pricing | None = None
per_request_limits: dict | None = None
top_provider: TopProvider | None = None
enabled: bool = True
upstream_provider_id: int | str | None = None
canonical_slug: str | None = None
alias_ids: list[str] | None = None
forwarded_model_id: str | None = None
reasoning: Reasoning | None = None
class Config:
extra = "ignore"
def __hash__(self) -> int:
return hash(self.id)
@validator("reasoning", pre=True)
def _coerce_reasoning(cls, value: object) -> object:
if value is None or value is False:
return None
if isinstance(value, Reasoning):
return None if value.is_empty() else value
if not isinstance(value, dict) or not value:
return None
try:
parsed = Reasoning.parse_obj(value)
except Exception:
return None
return None if parsed.is_empty() else parsed
def dict(self, **kwargs: object) -> dict:
# Non-reasoning models omit the field entirely so the catalog stays
# additive: existing clients never see a new null key.
data = super().dict(**kwargs) # type: ignore[arg-type]
reasoning = data.get("reasoning")
if not reasoning:
data.pop("reasoning", None)
elif isinstance(reasoning, dict):
cleaned = {k: v for k, v in reasoning.items() if v is not None}
if cleaned:
data["reasoning"] = cleaned
else:
data.pop("reasoning", None)
return data
def litellm_cost_entry(model_id: str) -> dict | None:
"""Look up ``model_id`` in litellm's bundled cost map.
litellm ships per-model USD rates keyed by the exact OpenRouter id
(``deepseek/deepseek-chat``) or the bare model name (``gpt-4o``,
``claude-sonnet-4-5``), so both spellings are tried. Keys are lowercase, so
a mixed-case upstream id (``deepseek-ai/DeepSeek-V4-Flash``) is retried via
a case-insensitive scan. Returns the matched cost dict, or ``None``.
"""
import litellm
candidates = (model_id, model_id.split("/", 1)[-1])
for key in candidates:
info = litellm.model_cost.get(key)
if isinstance(info, dict):
return info
lowered = {c.lower() for c in candidates}
for key, info in litellm.model_cost.items():
if isinstance(key, str) and key.lower() in lowered and isinstance(info, dict):
return info
return None
def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
"""Fill missing cache rates from litellm's bundled cost map.
The OpenRouter model feed omits ``input_cache_read``/``input_cache_write``
for many models (most DeepSeek entries, openai/gpt-4o, ...). Without a
cache rate, billing falls back to the full input rate, which overcharges
cache reads (DeepSeek hits are 10x cheaper) and undercharges Anthropic
cache writes (1.25x). The lookup (see ``litellm_cost_entry``) tries both
id spellings and a case-insensitive fallback.
Rates already present (e.g. provided by OpenRouter) are authoritative and
never overwritten. Unknown models are returned unchanged.
"""
needs_read = (pricing.input_cache_read or 0.0) <= 0.0
needs_write = (pricing.input_cache_write or 0.0) <= 0.0
if not (needs_read or needs_write):
return pricing
info = litellm_cost_entry(model_id)
if info is None:
return pricing
updated = Pricing.parse_obj(pricing.dict())
if needs_read:
read_rate = info.get("cache_read_input_token_cost")
if isinstance(read_rate, (int, float)) and read_rate > 0:
updated.input_cache_read = float(read_rate)
if needs_write:
write_rate = info.get("cache_creation_input_token_cost")
if isinstance(write_rate, (int, float)) and write_rate > 0:
updated.input_cache_write = float(write_rate)
return updated
def _has_valid_pricing(model: dict) -> bool:
"""Check if model has valid pricing (usable rates, and not free)."""
pricing = model.get("pricing", {})
if not pricing:
return False
# Coercion runs before the both-zero test below, which `NaN` would defeat
# on its own — and one entry the coercion chokes on must not unwind the
# whole fetch, which once cost the node an entire upstream catalog.
prompt = coerce_rate(pricing.get("prompt", 0))
completion = coerce_rate(pricing.get("completion", 0))
if prompt is None or completion is None:
return False
if prompt == 0 and completion == 0:
return False
return True
# OpenRouter occasionally answers /models with a truncated body, emptying the
# catalogue behind one log line. Retry, but keep 3 attempts within roughly the
# old single-attempt budget: this fetch blocks startup and the refresh loop.
OPENROUTER_MODELS_MAX_ATTEMPTS = 3
OPENROUTER_MODELS_TIMEOUT_SECONDS = 10
OPENROUTER_MODELS_RETRY_BACKOFF_SECONDS = 0.5
def _is_transient(error: BaseException) -> bool:
if isinstance(error, httpx.HTTPStatusError):
return error.response.status_code >= 500
return True
def _parse_models_response(response: httpx.Response | BaseException) -> list[dict]:
if isinstance(response, BaseException):
raise response
response.raise_for_status()
return [
model
for model in response.json().get("data", [])
if ":free" not in model.get("id", "").lower()
]
async def _fetch_openrouter_models_once(source_filter: str | None) -> list[dict]:
"""One attempt. Raises if /models is unusable; embeddings are best-effort."""
base_url = "https://openrouter.ai/api/v1"
timeout = OPENROUTER_MODELS_TIMEOUT_SECONDS
async with httpx.AsyncClient() as client:
models_response, embeddings_response = await asyncio.gather(
client.get(f"{base_url}/models", timeout=timeout),
client.get(f"{base_url}/embeddings/models", timeout=timeout),
return_exceptions=True,
)
# Losing /models is what empties the node, so it fails the attempt and
# the caller retries. A missing embeddings half must not do the same.
models_data = _parse_models_response(models_response)
try:
models_data.extend(_parse_models_response(embeddings_response))
except Exception as e:
logger.warning(f"Skipping OpenRouter embeddings models: {e}")
# Apply source filter and exclusions
filtered_models = []
for model in models_data:
model_id = model.get("id", "")
if source_filter:
source_prefix = f"{source_filter}/"
if not model_id.startswith(source_prefix):
continue
model = dict(model)
model["id"] = model_id[len(source_prefix) :]
model_id = model["id"]
if "(free)" in model.get("name", ""):
continue
if not _has_valid_pricing(model):
continue
filtered_models.append(model)
return filtered_models
async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
"""Fetch the OpenRouter catalogue; ``[]`` once every attempt has failed."""
for attempt in range(1, OPENROUTER_MODELS_MAX_ATTEMPTS + 1):
try:
return await _fetch_openrouter_models_once(source_filter)
except Exception as e:
last_attempt = attempt == OPENROUTER_MODELS_MAX_ATTEMPTS
if last_attempt or not _is_transient(e):
logger.error(
f"Error (async) fetching models from OpenRouter API "
f"after {attempt} attempt(s): {e}"
)
return []
logger.warning(
f"OpenRouter models fetch attempt {attempt}/"
f"{OPENROUTER_MODELS_MAX_ATTEMPTS} failed: {e}; retrying"
)
# Jittered so nodes do not retry in lockstep.
backoff = OPENROUTER_MODELS_RETRY_BACKOFF_SECONDS * attempt
await asyncio.sleep(backoff * random.uniform(0.5, 1.5))
return []
def _build_model_from_row(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
) -> Model:
"""The deterministic USD view of a stored model row, before the sats conversion."""
architecture = json.loads(row.architecture)
pricing = json.loads(row.pricing)
per_request_limits = (
json.loads(row.per_request_limits) if row.per_request_limits else None
)
top_provider_dict = json.loads(row.top_provider) if row.top_provider else None
# Rows written before the admin edge normalized rates can still carry
# numeric strings (``"0"``); compare as floats so the clamp cannot raise.
if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0:
pricing["request"] = 0.0
parsed_pricing = Pricing.parse_obj(pricing)
# Fill missing cache-read/write rates from litellm's cost map BEFORE applying
# the provider fee, so they carry the same markup as every other component.
# DB-stored override pricing (e.g. generic providers) omits cache rates;
# without this, ``_row_to_model`` bills cache reads at the full input rate —
# the ``_apply_provider_fee_to_model`` path backfills, but the override path
# used for admin-configured providers did not.
#
# Key on ``forwarded_model_id`` (the actual upstream model name litellm
# prices) when set: an alias row (id="local-alias",
# forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias
# and miss the cache rate.
pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id
parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing)
if apply_provider_fee:
parsed_pricing = Pricing.parse_obj(
{k: float(v) * provider_fee for k, v in parsed_pricing.dict().items()}
)
model = Model(
id=row.id,
name=row.name,
created=row.created,
description=row.description,
context_length=row.context_length,
architecture=Architecture.parse_obj(architecture),
pricing=parsed_pricing,
sats_pricing=None,
per_request_limits=per_request_limits,
top_provider=TopProvider.parse_obj(top_provider_dict)
if top_provider_dict
else None,
enabled=row.enabled,
upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None),
alias_ids=json.loads(row.alias_ids) if row.alias_ids else None,
forwarded_model_id=getattr(row, "forwarded_model_id", None),
)
if apply_provider_fee:
(
parsed_pricing.max_prompt_cost,
parsed_pricing.max_completion_cost,
parsed_pricing.max_cost,
) = _calculate_usd_max_costs(model)
return model
def _row_to_model(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
) -> Model:
model = _build_model_from_row(row, apply_provider_fee, provider_fee)
try:
sats_to_usd = sats_usd_price()
model = _update_model_sats_pricing(model, sats_to_usd)
except Exception as e:
logger.warning(f"Could not calculate sats pricing: {e}")
return model
async def list_models(
session: AsyncSession,
upstream_id: int,
include_disabled: bool = False,
apply_fees: bool = True,
) -> list[Model]:
from sqlmodel import select
from ..core.db import UpstreamProviderRow
query = select(ModelRow)
if upstream_id is not None:
query = query.where(ModelRow.upstream_provider_id == upstream_id)
if not include_disabled:
query = query.where(ModelRow.enabled)
rows = (await session.exec(query)).all() # type: ignore
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
models: list[Model] = []
for r in rows:
if not include_disabled and not (
r.upstream_provider_id in providers_by_id
and providers_by_id[r.upstream_provider_id].enabled
):
continue
try:
model = _row_to_model(
r,
apply_provider_fee=apply_fees,
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
if r.upstream_provider_id in providers_by_id
else 1.01,
)
except Exception as e:
# Stored pricing/architecture is JSON from whatever wrote the row, so
# a legacy import or foreign writer can leave a field that will not
# parse. Converting inside this loop meant one such row raised out of
# the whole listing and the node advertised nothing at all. Drop the
# row we cannot read — it is unservable either way — and keep serving
# the rest.
logger.warning(
"Skipping model row that could not be read",
extra={
"model_id": r.id,
"upstream_provider_id": r.upstream_provider_id,
"error": str(e),
"error_type": type(e).__name__,
},
)
continue
# Served-map backstop for legacy rows and writers that bypass the admin
# edge: a negative or non-finite rate is not a price. Serving one
# advertises a rate the cost calculation cannot bill on, so the request
# falls through to the flat maximum reservation — or, if the rate is
# negative, bills an amount settlement credits back to the caller.
# ``include_disabled`` is the operator's listing, which must keep showing
# the row so it can be repaired.
if not include_disabled and not has_usable_pricing(model.pricing):
logger.warning(
"Withholding model with an unusable stored rate from the catalog",
extra={
"model_id": r.id,
"upstream_provider_id": r.upstream_provider_id,
},
)
continue
models.append(model)
return models
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
"""Calculate max costs in USD based on model context/token limits.
Args:
model: Model object
Returns:
Tuple of (max_prompt_cost, max_completion_cost, max_cost) in USD
"""
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_usd = float(min_req_msat) / 1_000_000.0
prompt_price = model.pricing.prompt
completion_price = model.pricing.completion
if model.top_provider and (
model.top_provider.context_length or model.top_provider.max_completion_tokens
):
if (cl := model.top_provider.context_length) and (
mct := model.top_provider.max_completion_tokens
):
if cl <= mct:
return (
cl * prompt_price,
cl * completion_price,
cl * max(completion_price, prompt_price),
)
return (
cl * prompt_price,
mct * completion_price,
(cl - mct) * prompt_price + mct * completion_price,
)
elif cl := model.top_provider.context_length:
return (
cl * prompt_price,
cl * completion_price,
cl * max(completion_price, prompt_price),
)
elif mct := model.top_provider.max_completion_tokens:
return (
mct * prompt_price,
mct * completion_price,
mct * completion_price,
)
elif model.context_length:
return (
model.context_length * prompt_price,
model.context_length * completion_price,
model.context_length * max(completion_price, prompt_price),
)
p = prompt_price * 1_000_000
c = completion_price * 32_000
r = model.pricing.request * 100_000
i = model.pricing.image * 100
w = model.pricing.web_search * 1000
ir = model.pricing.internal_reasoning * 100
return (p, c, max(p + c + r + i + w + ir, min_req_usd))
def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
"""Update a model's sats_pricing based on USD pricing and exchange rate.
Args:
model: Model object to update
sats_to_usd: Current sats to USD exchange rate
Returns:
Updated Model object with new sats_pricing
"""
try:
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_sats = float(min_req_msat) / 1000.0
sats = Pricing.parse_obj(
{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
)
if sats.request <= 0.0:
sats.request = min_req_sats
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
return model.copy(update={"sats_pricing": sats})
except Exception as e:
logger.error(
"Failed to update sats pricing for model",
extra={
"model_id": model.id,
"error": str(e),
"error_type": type(e).__name__,
},
)
return model
async def _update_sats_pricing_once() -> None:
"""Update sats pricing once for all provider models (in-memory only)."""
from ..proxy import get_upstreams, refresh_model_maps
upstreams = get_upstreams()
if not upstreams:
return
sats_to_usd = sats_usd_price()
updated_count = 0
for upstream in upstreams:
updated_models = [
_update_model_sats_pricing(m, sats_to_usd)
for m in upstream.get_cached_models()
]
upstream._models_cache = updated_models
upstream._models_by_id = {
m.forwarded_model_id or m.id: m for m in updated_models
}
updated_count += len(updated_models)
if updated_count > 0:
logger.info(
f"Updated sats pricing for {updated_count} models",
extra={"models_updated": updated_count},
)
await refresh_model_maps()
async def update_sats_pricing() -> None:
"""Periodically update sats pricing for all provider models and database overrides."""
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
try:
await _update_sats_pricing_once()
except Exception as e:
logger.warning(
"Initial sats pricing update failed (will retry in loop)",
extra={"error": str(e)},
)
while True:
try:
interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
jitter = max(0.0, float(interval) * 0.1)
await asyncio.sleep(interval + random.uniform(0, jitter))
except asyncio.CancelledError:
break
try:
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
await _update_sats_pricing_once()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"Error updating sats pricing: {e}")
@models_router.get("/v1/models/paths")
@models_router.get("/v1/models/paths/", include_in_schema=False)
async def model_paths() -> dict:
"""All models with every upstream provider path they are reachable through."""
from ..upstream.model_paths import get_all_model_paths
return await get_all_model_paths()
@models_router.get("/v1/models/paths/model")
@models_router.get("/v1/models/paths/model/", include_in_schema=False)
async def model_paths_for_model(model_id: str) -> dict:
"""Paths for a single model.
Uses a query parameter (``?model_id=...``) under a fully static route so
model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL
encoding and there is no dynamic-route ambiguity.
"""
from ..proxy import get_unique_models
from ..upstream.model_paths import get_paths_for_model
result = await get_paths_for_model(model_id)
if not result["data"]:
advertised_ids = {model.id.lower() for model in get_unique_models()}
if model_id.lower() not in advertised_ids:
raise HTTPException(status_code=404, detail="Model not found")
return result
@models_router.get("/v1/models")
@models_router.get("/v1/models/", include_in_schema=False)
@models_router.get("/models")
@models_router.get("/models/", include_in_schema=False)
async def models(session: AsyncSession = Depends(get_session)) -> dict:
"""Get all available models from all providers with database overrides applied."""
from ..proxy import get_unique_models
items = get_unique_models()
data = []
for model in items:
data.append(model.dict())
return {"data": data}