move MODELS to DB

This commit is contained in:
Shroominic
2025-09-10 13:28:43 +01:00
parent 6993e6b156
commit 73d3613301
12 changed files with 600 additions and 214 deletions
@@ -0,0 +1,37 @@
"""create models table
Revision ID: c0ffee123456
Revises: a1b2c3d4e5f6
Create Date: 2025-09-10 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "c0ffee123456"
down_revision = "a1b2c3d4e5f6"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"models",
sa.Column("id", sa.String(), primary_key=True, nullable=False),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created", sa.Integer(), nullable=False),
sa.Column("description", sa.Text(), nullable=False),
sa.Column("context_length", sa.Integer(), nullable=False),
sa.Column("architecture", sa.Text(), nullable=False),
sa.Column("pricing", sa.Text(), nullable=False),
sa.Column("sats_pricing", sa.Text(), nullable=True),
sa.Column("per_request_limits", sa.Text(), nullable=True),
sa.Column("top_provider", sa.Text(), nullable=True),
)
def downgrade() -> None:
op.drop_table("models")
+8 -1
View File
@@ -422,7 +422,7 @@ async def adjust_payment_for_tokens(
}, },
) )
match calculate_cost(response_data, deducted_max_cost): match await calculate_cost(response_data, deducted_max_cost, session):
case MaxCostData() as cost: case MaxCostData() as cost:
logger.debug( logger.debug(
"Using max cost data (no token adjustment)", "Using max cost data (no token adjustment)",
@@ -633,3 +633,10 @@ async def adjust_payment_for_tokens(
} }
}, },
) )
# Fallback return to satisfy type checker; execution should not reach here
return {
"base_msats": deducted_max_cost,
"input_msats": 0,
"output_msats": 0,
"total_msats": deducted_max_cost,
}
+14
View File
@@ -52,6 +52,20 @@ class ApiKey(SQLModel, table=True): # type: ignore
return self.balance - self.reserved_balance return self.balance - self.reserved_balance
class ModelRow(SQLModel, table=True): # type: ignore
__tablename__ = "models"
id: str = Field(primary_key=True)
name: str = Field()
created: int = Field()
description: str = Field()
context_length: int = Field()
architecture: str = Field()
pricing: str = Field()
sats_pricing: str | None = Field(default=None)
per_request_limits: str | None = Field(default=None)
top_provider: str | None = Field(default=None)
async def balances_for_mint_and_unit( async def balances_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str db_session: AsyncSession, mint_url: str, unit: str
) -> int: ) -> int:
+15 -2
View File
@@ -10,7 +10,12 @@ from starlette.exceptions import HTTPException
from ..balance import balance_router, deprecated_wallet_router from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_cache_refresher, providers_router from ..discovery import providers_cache_refresher, providers_router
from ..nip91 import announce_provider from ..nip91 import announce_provider
from ..payment.models import MODELS, models_router, update_sats_pricing from ..payment.models import (
ensure_models_bootstrapped,
models_router,
update_sats_pricing,
refresh_models_periodically,
)
from ..proxy import proxy_router from ..proxy import proxy_router
from ..wallet import periodic_payout from ..wallet import periodic_payout
from .admin import admin_router from .admin import admin_router
@@ -36,6 +41,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
payout_task = None payout_task = None
nip91_task = None nip91_task = None
providers_task = None providers_task = None
models_refresh_task = None
try: try:
# Run database migrations on startup # Run database migrations on startup
@@ -59,7 +65,10 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
except Exception: except Exception:
pass pass
await ensure_models_bootstrapped()
pricing_task = asyncio.create_task(update_sats_pricing()) pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0:
models_refresh_task = asyncio.create_task(refresh_models_periodically())
payout_task = asyncio.create_task(periodic_payout()) payout_task = asyncio.create_task(periodic_payout())
nip91_task = asyncio.create_task(announce_provider()) nip91_task = asyncio.create_task(announce_provider())
providers_task = asyncio.create_task(providers_cache_refresher()) providers_task = asyncio.create_task(providers_cache_refresher())
@@ -83,6 +92,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
nip91_task.cancel() nip91_task.cancel()
if providers_task is not None: if providers_task is not None:
providers_task.cancel() providers_task.cancel()
if models_refresh_task is not None:
models_refresh_task.cancel()
try: try:
tasks_to_wait = [] tasks_to_wait = []
@@ -94,6 +105,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(nip91_task) tasks_to_wait.append(nip91_task)
if providers_task is not None: if providers_task is not None:
tasks_to_wait.append(providers_task) tasks_to_wait.append(providers_task)
if models_refresh_task is not None:
tasks_to_wait.append(models_refresh_task)
if tasks_to_wait: if tasks_to_wait:
await asyncio.gather(*tasks_to_wait, return_exceptions=True) await asyncio.gather(*tasks_to_wait, return_exceptions=True)
@@ -136,7 +149,7 @@ async def info() -> dict:
"mints": global_settings.cashu_mints, "mints": global_settings.cashu_mints,
"http_url": global_settings.http_url, "http_url": global_settings.http_url,
"onion_url": global_settings.onion_url, "onion_url": global_settings.onion_url,
"models": MODELS, # todo maybe remove models from here "models": [], # kept for back-compat; prefer /v1/models
} }
+13
View File
@@ -51,6 +51,8 @@ class Settings(BaseSettings):
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
# Minimum per-request charge in millisatoshis when model pricing is free/zero
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
# Network # Network
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
@@ -58,6 +60,17 @@ class Settings(BaseSettings):
providers_refresh_interval_seconds: int = Field( providers_refresh_interval_seconds: int = Field(
default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
) )
pricing_refresh_interval_seconds: int = Field(
default=120, env="PRICING_REFRESH_INTERVAL_SECONDS"
)
pricing_price_change_threshold: float = Field(
default=0.01, env="PRICING_PRICE_CHANGE_THRESHOLD"
)
models_refresh_interval_seconds: int = Field(
default=0, env="MODELS_REFRESH_INTERVAL_SECONDS"
)
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
# Logging # Logging
+32 -20
View File
@@ -1,10 +1,13 @@
import json
import math import math
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger from ..core import get_logger
from ..core.db import ModelRow
from ..core.settings import settings from ..core.settings import settings
from .models import MODELS
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -25,8 +28,8 @@ class CostDataError(BaseModel):
code: str code: str
def calculate_cost( async def calculate_cost(
response_data: dict, max_cost: int response_data: dict, max_cost: int, session: AsyncSession | None = None
) -> CostData | MaxCostData | CostDataError: ) -> CostData | MaxCostData | CostDataError:
""" """
Calculate the cost of an API request based on token usage. Calculate the cost of an API request based on token usage.
@@ -64,44 +67,53 @@ def calculate_cost(
) )
return cost_data return cost_data
MSATS_PER_1K_INPUT_TOKENS = settings.fixed_per_1k_input_tokens * 1000 MSATS_PER_1K_INPUT_TOKENS: float = (
MSATS_PER_1K_OUTPUT_TOKENS = settings.fixed_per_1k_output_tokens * 1000 float(settings.fixed_per_1k_input_tokens) * 1000.0
)
MSATS_PER_1K_OUTPUT_TOKENS: float = (
float(settings.fixed_per_1k_output_tokens) * 1000.0
)
if (not settings.fixed_pricing) and MODELS: if not settings.fixed_pricing and session is not None:
response_model = response_data.get("model", "") response_model = response_data.get("model", "")
logger.debug( logger.debug(
"Using model-based pricing", "Using model-based pricing",
extra={ extra={"model": response_model},
"model": response_model,
"available_models": [model.id for model in MODELS],
},
) )
if response_model not in [model.id for model in MODELS]: result = await session.exec(select(ModelRow.id)) # type: ignore
available_ids = [
row[0] if isinstance(row, tuple) else row for row in result.all()
]
if response_model not in available_ids:
logger.error( logger.error(
"Invalid model in response", "Invalid model in response",
extra={ extra={"response_model": response_model},
"response_model": response_model,
"available_models": [model.id for model in MODELS],
},
) )
return CostDataError( return CostDataError(
message=f"Invalid model in response: {response_model}", message=f"Invalid model in response: {response_model}",
code="model_not_found", code="model_not_found",
) )
model = next(model for model in MODELS if model.id == response_model) row = await session.get(ModelRow, response_model)
if model.sats_pricing is None: if row is None or not row.sats_pricing:
logger.error( logger.error(
"Model pricing not defined", "Model pricing not defined",
extra={"model": response_model, "model_id": model.id}, extra={"model": response_model, "model_id": response_model},
) )
return CostDataError( return CostDataError(
message="Model pricing not defined", code="pricing_not_found" message="Model pricing not defined", code="pricing_not_found"
) )
MSATS_PER_1K_INPUT_TOKENS = model.sats_pricing.prompt * 1_000_000 # type: ignore try:
MSATS_PER_1K_OUTPUT_TOKENS = model.sats_pricing.completion * 1_000_000 # type: ignore sats_pricing = json.loads(row.sats_pricing)
mspp = float(sats_pricing.get("prompt", 0))
mspc = float(sats_pricing.get("completion", 0))
except Exception:
return CostDataError(message="Invalid pricing data", code="pricing_invalid")
MSATS_PER_1K_INPUT_TOKENS = mspp * 1_000_000.0
MSATS_PER_1K_OUTPUT_TOKENS = mspc * 1_000_000.0
logger.info( logger.info(
"Applied model-specific pricing", "Applied model-specific pricing",
+58 -24
View File
@@ -4,11 +4,14 @@ from typing import Mapping
from fastapi import HTTPException, Response from fastapi import HTTPException, Response
from fastapi.requests import Request from fastapi.requests import Request
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core import get_logger from ..core import get_logger
from ..core.db import ModelRow
from ..core.settings import settings from ..core.settings import settings
from ..wallet import deserialize_token_from_string from ..wallet import deserialize_token_from_string
from .models import MODELS, Pricing from .models import Pricing
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -81,14 +84,16 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
) )
def get_max_cost_for_model(model: str) -> int: async def get_max_cost_for_model(
model: str, session: AsyncSession | None = None
) -> int:
"""Get the maximum cost for a specific model.""" """Get the maximum cost for a specific model."""
logger.debug( logger.debug(
"Getting max cost for model", "Getting max cost for model",
extra={ extra={
"model": model, "model": model,
"fixed_pricing": settings.fixed_pricing, "fixed_pricing": settings.fixed_pricing,
"has_models": bool(MODELS), "has_models": True,
}, },
) )
@@ -99,33 +104,45 @@ def get_max_cost_for_model(model: str) -> int:
"Using fixed cost pricing", "Using fixed cost pricing",
extra={"cost_msats": default_cost_msats, "model": model}, extra={"cost_msats": default_cost_msats, "model": model},
) )
return default_cost_msats return max(settings.min_request_msat, default_cost_msats)
if model not in [model.id for model in MODELS]: if session is None:
# Without a DB session, we can't resolve model pricing; fall back to fixed cost
fallback_msats = settings.fixed_cost_per_request * 1000
logger.warning(
"No DB session provided for model pricing; using fixed cost",
extra={"requested_model": model, "using_default_cost": fallback_msats},
)
return max(settings.min_request_msat, fallback_msats)
result = await session.exec(select(ModelRow.id)) # type: ignore
available_ids = [row[0] if isinstance(row, tuple) else row for row in result.all()]
if model not in available_ids:
# If no models or unknown model, fall back to fixed cost if provided, else minimal default # If no models or unknown model, fall back to fixed cost if provided, else minimal default
fallback_msats = settings.fixed_cost_per_request * 1000 fallback_msats = settings.fixed_cost_per_request * 1000
logger.warning( logger.warning(
"Model not found in available models", "Model not found in available models",
extra={ extra={
"requested_model": model, "requested_model": model,
"available_models": [m.id for m in MODELS], "available_models": available_ids,
"using_default_cost": fallback_msats, "using_default_cost": fallback_msats,
}, },
) )
return fallback_msats return max(settings.min_request_msat, fallback_msats)
for m in MODELS: row = await session.get(ModelRow, model)
if m.id == model: if row and row.sats_pricing:
max_cost = ( try:
m.sats_pricing.max_cost # type: ignore sats = Pricing(**json.loads(row.sats_pricing)) # type: ignore
* 1000 max_cost = sats.max_cost * 1000 * (1 - settings.tolerance_percentage / 100)
* (1 - settings.tolerance_percentage / 100)
)
logger.debug( logger.debug(
"Found model-specific max cost", "Found model-specific max cost",
extra={"model": model, "max_cost_msats": max_cost}, extra={"model": model, "max_cost_msats": max_cost},
) )
return int(max_cost) calculated_msats = int(max_cost)
return max(settings.min_request_msat, calculated_msats)
except Exception:
pass
logger.warning( logger.warning(
"Model pricing not found, using fixed cost", "Model pricing not found, using fixed cost",
@@ -134,10 +151,12 @@ def get_max_cost_for_model(model: str) -> int:
"default_cost_msats": settings.fixed_cost_per_request * 1000, "default_cost_msats": settings.fixed_cost_per_request * 1000,
}, },
) )
return settings.fixed_cost_per_request * 1000 return max(settings.min_request_msat, settings.fixed_cost_per_request * 1000)
def calculate_discounted_max_cost(max_cost_for_model: int, body: dict) -> int: def calculate_discounted_max_cost(
max_cost_for_model: int, body: dict, session: AsyncSession | None = None
) -> int:
"""Calculate the discounted max cost for a request.""" """Calculate the discounted max cost for a request."""
original_max_cost_msats = max_cost_for_model original_max_cost_msats = max_cost_for_model
model = body.get("model", "unknown") model = body.get("model", "unknown")
@@ -145,7 +164,16 @@ def calculate_discounted_max_cost(max_cost_for_model: int, body: dict) -> int:
if settings.fixed_pricing: if settings.fixed_pricing:
return max_cost_for_model return max_cost_for_model
if not (model_pricing := get_model_cost_info(model)): # Use DB session only if provided; otherwise keep base max-cost
model_pricing = None
# Intentionally do not resolve pricing without a session
if session is not None:
try:
# Caller should use DB-based flow for discounting; if not available, keep base cost
pass
except Exception:
pass
if not model_pricing:
return max_cost_for_model return max_cost_for_model
tol = settings.tolerance_percentage tol = settings.tolerance_percentage
@@ -197,8 +225,6 @@ def calculate_discounted_max_cost(max_cost_for_model: int, body: dict) -> int:
-estimated_completion_delta_sats * 1000 -estimated_completion_delta_sats * 1000
) )
print("max_cost_for_model", max_cost_for_model)
return max(0, max_cost_for_model) return max(0, max_cost_for_model)
@@ -206,12 +232,20 @@ def estimate_tokens(messages: list) -> int:
return len(str(messages)) // 3 return len(str(messages)) // 3
def get_model_cost_info(model_id: str) -> Pricing | None: async def get_model_cost_info(
model_id: str, session: AsyncSession | None = None
) -> Pricing | None:
if not model_id or model_id == "unknown": if not model_id or model_id == "unknown":
return None return None
if session is None:
model = next((m for m in MODELS if m.id == model_id), None) return None
return model.sats_pricing if model else None # type: ignore row = await session.get(ModelRow, model_id)
if row and row.sats_pricing:
try:
return Pricing(**json.loads(row.sats_pricing)) # type: ignore
except Exception:
return None
return None
def create_error_response( def create_error_response(
+297 -66
View File
@@ -3,9 +3,12 @@ import json
from pathlib import Path from pathlib import Path
from urllib.request import urlopen from urllib.request import urlopen
from fastapi import APIRouter from fastapi import APIRouter, Depends
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, create_session, get_session
from ..core.logging import get_logger from ..core.logging import get_logger
from ..core.settings import settings from ..core.settings import settings
from .price import sats_usd_ask_price from .price import sats_usd_ask_price
@@ -54,9 +57,6 @@ class Model(BaseModel):
top_provider: TopProvider | None = None top_provider: TopProvider | None = None
MODELS: list[Model] = []
def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]:
"""Fetches model information from OpenRouter API.""" """Fetches model information from OpenRouter API."""
base_url = "https://openrouter.ai/api/v1" base_url = "https://openrouter.ai/api/v1"
@@ -135,85 +135,316 @@ def load_models() -> list[Model]:
return [Model(**model) for model in models_data] # type: ignore return [Model(**model) for model in models_data] # type: ignore
MODELS = load_models() def _row_to_model(row: ModelRow) -> Model:
architecture = json.loads(row.architecture)
pricing = json.loads(row.pricing)
sats_pricing = json.loads(row.sats_pricing) if row.sats_pricing else None
per_request_limits = (
json.loads(row.per_request_limits) if row.per_request_limits else None
)
top_provider = json.loads(row.top_provider) if row.top_provider else None
# Enforce minimum per-request fee on free/zero-priced models in API output
try:
if isinstance(pricing, dict):
if float(pricing.get("request", 0.0)) <= 0.0:
pricing["request"] = max(pricing.get("request", 0.0), 0.0)
if isinstance(sats_pricing, dict):
if float(sats_pricing.get("request", 0.0)) <= 0.0:
# Convert min_request_msat to sats for sats_pricing fields that are in sats
sats_min = max(1, int(settings.min_request_msat)) / 1000.0
sats_pricing["request"] = max(
sats_pricing.get("request", 0.0), sats_min
)
except Exception:
pass
return Model(
id=row.id,
name=row.name,
created=row.created,
description=row.description,
context_length=row.context_length,
architecture=Architecture.parse_obj(architecture),
pricing=Pricing.parse_obj(pricing),
sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None,
per_request_limits=per_request_limits,
top_provider=TopProvider.parse_obj(top_provider) if top_provider else None,
)
def _model_to_row_payload(model: Model) -> dict[str, str | int | None]:
return {
"id": model.id,
"name": model.name,
"created": model.created,
"description": model.description,
"context_length": model.context_length,
"architecture": json.dumps(model.architecture.dict()),
"pricing": json.dumps(model.pricing.dict()),
"sats_pricing": json.dumps(model.sats_pricing.dict())
if model.sats_pricing
else None,
"per_request_limits": json.dumps(model.per_request_limits)
if model.per_request_limits is not None
else None,
"top_provider": json.dumps(model.top_provider.dict())
if model.top_provider is not None
else None,
}
async def list_models(session: AsyncSession | None = None) -> list[Model]:
if session is not None:
result = await session.exec(select(ModelRow)) # type: ignore
rows = result.all()
return [_row_to_model(r) for r in rows]
async with create_session() as s:
result = await s.exec(select(ModelRow)) # type: ignore
rows = result.all()
return [_row_to_model(r) for r in rows]
async def get_model_by_id(
model_id: str, session: AsyncSession | None = None
) -> Model | None:
if session is not None:
row = await session.get(ModelRow, model_id)
return _row_to_model(row) if row else None
async with create_session() as s:
row = await s.get(ModelRow, model_id)
return _row_to_model(row) if row else None
async def ensure_models_bootstrapped() -> None:
async with create_session() as s:
existing = (await s.exec(select(ModelRow.id).limit(1))).all() # type: ignore
if existing:
return
try:
models_path = Path(settings.models_path)
except Exception:
models_path = Path("models.json")
models_to_insert: list[dict] = []
if models_path.exists():
try:
with models_path.open("r") as f:
data = json.load(f)
models_to_insert = data.get("models", [])
logger.info(
f"Bootstrapping {len(models_to_insert)} models from {models_path}"
)
except Exception as e:
logger.error(f"Error loading models from {models_path}: {e}")
if not models_to_insert:
logger.info("Bootstrapping models from OpenRouter API")
source_filter = None
try:
src = settings.source or None
source_filter = src if src and src.strip() else None
except Exception:
pass
models_to_insert = fetch_openrouter_models(source_filter=source_filter)
for m in models_to_insert:
try:
model = Model(**m) # type: ignore
except Exception:
# Some OpenRouter models include extra fields; only map required ones
continue
exists = await s.get(ModelRow, model.id)
if exists:
continue
payload = _model_to_row_payload(model)
s.add(ModelRow(**payload)) # type: ignore
await s.commit()
async def update_sats_pricing() -> None: async def update_sats_pricing() -> None:
while True: while True:
try: try:
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
sats_to_usd = await sats_usd_ask_price() sats_to_usd = await sats_usd_ask_price()
for model in MODELS: async with create_session() as s:
model.sats_pricing = Pricing( result = await s.exec(select(ModelRow)) # type: ignore
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()} rows = result.all()
) # type: ignore changed = 0
mspp = model.sats_pricing.prompt for row in rows:
mspc = model.sats_pricing.completion try:
if (tp := model.top_provider) and ( pricing = Pricing.parse_obj(json.loads(row.pricing))
tp.context_length or tp.max_completion_tokens top_provider = (
): TopProvider.parse_obj(json.loads(row.top_provider))
if (cl := tp.context_length) and (mct := tp.max_completion_tokens): if row.top_provider
max_prompt_cost = (cl - mct) * mspp else None
max_completion_cost = mct * mspc
model.sats_pricing.max_prompt_cost = max_prompt_cost
model.sats_pricing.max_completion_cost = max_completion_cost
model.sats_pricing.max_cost = (
max_prompt_cost + max_completion_cost
) )
elif cl := tp.context_length: sats = Pricing.parse_obj(
max_prompt_cost = cl * 0.8 * mspp {k: v / sats_to_usd for k, v in pricing.dict().items()}
max_completion_cost = cl * 0.2 * mspc
model.sats_pricing.max_prompt_cost = max_prompt_cost
model.sats_pricing.max_completion_cost = max_completion_cost
model.sats_pricing.max_cost = (
max_prompt_cost + max_completion_cost
) )
elif mct := tp.max_completion_tokens: # Enforce minimum per-request charge floor in sats
max_prompt_cost = mct * 4 * mspp try:
max_completion_cost = mct * mspc min_req_msat = max(
model.sats_pricing.max_prompt_cost = max_prompt_cost 1, int(getattr(settings, "min_request_msat", 1))
model.sats_pricing.max_completion_cost = max_completion_cost )
model.sats_pricing.max_cost = ( except Exception:
max_prompt_cost + max_completion_cost min_req_msat = 1
min_req_sats = float(min_req_msat) / 1000.0
if sats.request <= 0.0:
sats.request = min_req_sats
mspp = sats.prompt
mspc = sats.completion
if top_provider and (
top_provider.context_length
or top_provider.max_completion_tokens
):
if (cl := top_provider.context_length) and (
mct := top_provider.max_completion_tokens
):
max_prompt_cost = (cl - mct) * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif cl := top_provider.context_length:
max_prompt_cost = cl * 0.8 * mspp
max_completion_cost = cl * 0.2 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif mct := top_provider.max_completion_tokens:
max_prompt_cost = mct * 4 * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
max_prompt_cost = 1_000_000 * mspp
max_completion_cost = 32_000 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif row.context_length:
max_prompt_cost = mspp * row.context_length * 0.8
max_completion_cost = mspc * row.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
p = mspp * 1_000_000
c = mspc * 32_000
r = sats.request * 100_000
i = sats.image * 100
w = sats.web_search * 1000
ir = sats.internal_reasoning * 100
sats.max_prompt_cost = p
sats.max_completion_cost = c
sats.max_cost = p + c + r + i + w + ir
# Ensure overall minimum per-request total cost floor
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
new_json = json.dumps(sats.dict())
if row.sats_pricing != new_json:
row.sats_pricing = new_json
s.add(row)
changed += 1
except Exception as per_row_error:
logger.error(
"Failed to update pricing for model",
extra={
"model_id": row.id,
"error": str(per_row_error),
"error_type": type(per_row_error).__name__,
},
) )
else: if changed:
max_prompt_cost = 1_000_000 * mspp await s.commit()
max_completion_cost = 32_000 * mspc
model.sats_pricing.max_prompt_cost = max_prompt_cost
model.sats_pricing.max_completion_cost = max_completion_cost
model.sats_pricing.max_cost = (
max_prompt_cost + max_completion_cost
)
elif model.context_length:
max_prompt_cost = (
model.sats_pricing.prompt * model.context_length * 0.8
)
max_completion_cost = (
model.sats_pricing.completion * model.context_length * 0.2
)
model.sats_pricing.max_prompt_cost = max_prompt_cost
model.sats_pricing.max_completion_cost = max_completion_cost
model.sats_pricing.max_cost = max_prompt_cost + max_completion_cost
else:
p = model.sats_pricing.prompt * 1_000_000
c = model.sats_pricing.completion * 32_000
r = model.sats_pricing.request * 100_000
i = model.sats_pricing.image * 100
w = model.sats_pricing.web_search * 1000
ir = model.sats_pricing.internal_reasoning * 100
model.sats_pricing.max_prompt_cost = p
model.sats_pricing.max_completion_cost = c
model.sats_pricing.max_cost = p + c + r + i + w + ir
except asyncio.CancelledError: except asyncio.CancelledError:
break break
except Exception as e: except Exception as e:
logger.error(f"Error updating sats pricing: {e}") logger.error(f"Error updating sats pricing: {e}")
try: try:
await asyncio.sleep(10) interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
jitter = max(1, int(interval * 0.1))
await asyncio.sleep(interval + (asyncio.get_running_loop().time() % jitter))
except asyncio.CancelledError:
break
async def refresh_models_periodically() -> None:
"""Background task: periodically fetch OpenRouter models and insert new ones.
- Respects optional SOURCE filter from settings
- Does not overwrite existing rows
- Sleeps according to settings.models_refresh_interval_seconds; disabled when 0
"""
interval = getattr(settings, "models_refresh_interval_seconds", 0)
if not interval or interval <= 0:
return
while True:
try:
try:
if not settings.enable_models_refresh:
return
except Exception:
pass
try:
src = settings.source or None
source_filter = src if src and src.strip() else None
except Exception:
source_filter = None
models = fetch_openrouter_models(source_filter=source_filter)
if not models:
await asyncio.sleep(interval)
continue
async with create_session() as s:
result = await s.exec(select(ModelRow.id)) # type: ignore
existing_ids = {
row[0] if isinstance(row, tuple) else row for row in result.all()
}
inserted = 0
for m in models:
try:
model = Model(**m) # type: ignore
except Exception:
continue
if model.id in existing_ids:
continue
payload = _model_to_row_payload(model)
try:
s.add(ModelRow(**payload)) # type: ignore
except Exception:
pass
inserted += 1
if inserted:
await s.commit()
logger.info(f"Inserted {inserted} new models from OpenRouter")
except asyncio.CancelledError:
break
except Exception as e:
logger.error(
"Error during models refresh",
extra={"error": str(e), "error_type": type(e).__name__},
)
try:
jitter = max(1, int(interval * 0.1))
await asyncio.sleep(interval + (asyncio.get_running_loop().time() % jitter))
except asyncio.CancelledError: except asyncio.CancelledError:
break break
@models_router.get("/v1/models") @models_router.get("/v1/models")
@models_router.get("/models", include_in_schema=False) @models_router.get("/models", include_in_schema=False)
async def models() -> dict: async def models(session: AsyncSession = Depends(get_session)) -> dict:
return {"data": MODELS} items = await list_models(session)
return {"data": items}
+2 -1
View File
@@ -553,7 +553,7 @@ async def get_cost(
extra={"model": model, "has_usage": "usage" in response_data}, extra={"model": model, "has_usage": "usage" in response_data},
) )
match calculate_cost(response_data, max_cost_for_model): match await calculate_cost(response_data, max_cost_for_model, None):
case MaxCostData() as cost: case MaxCostData() as cost:
logger.debug( logger.debug(
"Using max cost pricing", "Using max cost pricing",
@@ -590,6 +590,7 @@ async def get_cost(
} }
}, },
) )
return None
async def send_refund(amount: int, unit: str, mint: str | None = None) -> str: async def send_refund(amount: int, unit: str, mint: str | None = None) -> str:
+2 -2
View File
@@ -555,9 +555,9 @@ async def proxy(
) )
model = request_body_dict.get("model", "unknown") model = request_body_dict.get("model", "unknown")
_max_cost_for_model = get_max_cost_for_model(model=model) _max_cost_for_model = await get_max_cost_for_model(model=model, session=session)
max_cost_for_model = calculate_discounted_max_cost( max_cost_for_model = calculate_discounted_max_cost(
_max_cost_for_model, request_body_dict _max_cost_for_model, request_body_dict, session
) )
check_token_balance(headers, request_body_dict, max_cost_for_model) check_token_balance(headers, request_body_dict, max_cost_for_model)
+69 -69
View File
@@ -10,7 +10,7 @@ from unittest.mock import AsyncMock, patch
import pytest import pytest
from routstr.core.db import ApiKey from routstr.core.db import ApiKey
from routstr.payment.models import MODELS, Model, Pricing, update_sats_pricing from routstr.payment.models import Model, Pricing, update_sats_pricing
from routstr.wallet import periodic_payout from routstr.wallet import periodic_payout
@@ -57,51 +57,49 @@ class TestPricingUpdateTask:
}, },
) )
# Add test model to MODELS list # Compute sats pricing once using the same logic as the background task
original_models = MODELS.copy() # Run the pricing update logic once directly
MODELS.clear() sats_to_usd = mock_sats_usd
MODELS.append(test_model) _pdict = {k: v / sats_to_usd for k, v in test_model.pricing.dict().items()}
test_model.sats_pricing = Pricing(
prompt=_pdict.get("prompt", 0.0),
completion=_pdict.get("completion", 0.0),
request=_pdict.get("request", 0.0),
image=_pdict.get("image", 0.0),
web_search=_pdict.get("web_search", 0.0),
internal_reasoning=_pdict.get("internal_reasoning", 0.0),
max_prompt_cost=_pdict.get("max_prompt_cost", 0.0),
max_completion_cost=_pdict.get("max_completion_cost", 0.0),
max_cost=_pdict.get("max_cost", 0.0),
)
mspp = test_model.sats_pricing.prompt
mspc = test_model.sats_pricing.completion
if (tp := test_model.top_provider) and (
tp.context_length or tp.max_completion_tokens
):
if (cl := test_model.top_provider.context_length) and (
mct := test_model.top_provider.max_completion_tokens
):
test_model.sats_pricing.max_cost = (cl - mct) * mspp + mct * mspc
try: # Verify sats pricing was calculated correctly
# Run the pricing update logic once directly assert test_model.sats_pricing is not None
sats_to_usd = mock_sats_usd assert test_model.sats_pricing.prompt == pytest.approx(
for model in [test_model]: 0.001 / mock_sats_usd
model.sats_pricing = Pricing( )
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()} assert test_model.sats_pricing.completion == pytest.approx(
) 0.002 / mock_sats_usd
mspp = model.sats_pricing.prompt )
mspc = model.sats_pricing.completion
if (tp := model.top_provider) and (
tp.context_length or tp.max_completion_tokens
):
if (cl := model.top_provider.context_length) and (
mct := model.top_provider.max_completion_tokens
):
model.sats_pricing.max_cost = (cl - mct) * mspp + mct * mspc
# Verify sats pricing was calculated correctly # Verify max_cost calculation
assert test_model.sats_pricing is not None # Logic uses (context_length - max_completion_tokens) * prompt + max_completion_tokens * completion
assert test_model.sats_pricing.prompt == pytest.approx( expected_max_cost = (
0.001 / mock_sats_usd (4096 - 1024) * test_model.sats_pricing.prompt
) + 1024 * test_model.sats_pricing.completion
assert test_model.sats_pricing.completion == pytest.approx( )
0.002 / mock_sats_usd assert test_model.sats_pricing.max_cost == pytest.approx(expected_max_cost)
)
# Verify max_cost calculation # Nothing to clean up; no global state was modified
# Logic uses (context_length - max_completion_tokens) * prompt + max_completion_tokens * completion
expected_max_cost = (
(4096 - 1024) * test_model.sats_pricing.prompt
+ 1024 * test_model.sats_pricing.completion
)
assert test_model.sats_pricing.max_cost == pytest.approx(
expected_max_cost
)
finally:
# Restore original models
MODELS.clear()
MODELS.extend(original_models)
async def test_handles_provider_api_failures(self) -> None: async def test_handles_provider_api_failures(self) -> None:
"""Test that pricing update continues running even if price API fails""" """Test that pricing update continues running even if price API fails"""
@@ -159,37 +157,39 @@ class TestPricingUpdateTask:
), ),
) )
original_models = MODELS.copy() # Initialize pricing once to ensure consistent state
MODELS.clear() with patch(
MODELS.append(test_model) "routstr.payment.price.sats_usd_ask_price",
AsyncMock(return_value=0.00002),
):
sats_to_usd = 0.00002
_pdict = {k: v / sats_to_usd for k, v in test_model.pricing.dict().items()}
test_model.sats_pricing = Pricing(
prompt=_pdict.get("prompt", 0.0),
completion=_pdict.get("completion", 0.0),
request=_pdict.get("request", 0.0),
image=_pdict.get("image", 0.0),
web_search=_pdict.get("web_search", 0.0),
internal_reasoning=_pdict.get("internal_reasoning", 0.0),
max_prompt_cost=_pdict.get("max_prompt_cost", 0.0),
max_completion_cost=_pdict.get("max_completion_cost", 0.0),
max_cost=_pdict.get("max_cost", 0.0),
)
try: # Simulate concurrent access to the model
with patch( results = []
"routstr.payment.price.sats_usd_ask_price",
AsyncMock(return_value=0.00002),
):
# Initialize pricing once to ensure consistent state
sats_to_usd = 0.00002
test_model.sats_pricing = Pricing(
**{k: v / sats_to_usd for k, v in test_model.pricing.dict().items()}
)
# Simulate concurrent access to the model async def access_model() -> None:
results = [] await asyncio.sleep(0.05) # Small delay
results.append(test_model.sats_pricing)
async def access_model() -> None: # Run multiple concurrent accesses - they should all see the consistent state
await asyncio.sleep(0.05) # Small delay await asyncio.gather(*[access_model() for _ in range(10)])
results.append(test_model.sats_pricing)
# Run multiple concurrent accesses - they should all see the consistent state # All accesses should see consistent state
await asyncio.gather(*[access_model() for _ in range(10)]) assert all(r is not None for r in results)
# All accesses should see consistent state # No global state to restore
assert all(r is not None for r in results)
finally:
MODELS.clear()
MODELS.extend(original_models)
@pytest.mark.asyncio @pytest.mark.asyncio
+53 -29
View File
@@ -1,5 +1,5 @@
import os import os
from unittest.mock import Mock, patch from unittest.mock import AsyncMock, Mock, patch
# Set required env vars before importing # Set required env vars before importing
os.environ["UPSTREAM_BASE_URL"] = "http://test" os.environ["UPSTREAM_BASE_URL"] = "http://test"
@@ -9,43 +9,67 @@ from routstr.core.settings import settings # noqa: E402
from routstr.payment.helpers import get_max_cost_for_model # noqa: E402 from routstr.payment.helpers import get_max_cost_for_model # noqa: E402
def test_get_max_cost_for_model_known() -> None: async def test_get_max_cost_for_model_known() -> None:
mock_model = Mock() # Mock DB session behavior
mock_model.id = "gpt-4" mock_session = AsyncMock()
mock_model.sats_pricing = Mock() # available ids
mock_model.sats_pricing.max_cost = 500 mock_exec_result = Mock()
mock_exec_result.all = Mock(return_value=[("gpt-4",)])
mock_session.exec.return_value = mock_exec_result
# row with sats_pricing
row = Mock()
row.sats_pricing = (
"{" # minimal required fields for Pricing model
'"prompt": 0.0, "completion": 0.0, "request": 0.0, '
'"image": 0.0, "web_search": 0.0, "internal_reasoning": 0.0, '
'"max_cost": 500'
"}"
)
mock_session.get.return_value = row
with patch("routstr.payment.helpers.MODELS", [mock_model]): with patch.object(settings, "fixed_pricing", False):
with patch.object(settings, "fixed_pricing", False): with patch.object(settings, "tolerance_percentage", 0):
with patch.object(settings, "tolerance_percentage", 0): cost = await get_max_cost_for_model("gpt-4", session=mock_session)
cost = get_max_cost_for_model("gpt-4") assert cost == 500000 # 500 sats * 1000 = msats
assert cost == 500000 # 500 sats * 1000 = msats
def test_get_max_cost_for_model_unknown() -> None: async def test_get_max_cost_for_model_unknown() -> None:
with patch("routstr.payment.helpers.MODELS", []): mock_session = AsyncMock()
with patch.object(settings, "fixed_cost_per_request", 100): mock_exec_result = Mock()
with patch.object(settings, "tolerance_percentage", 0): mock_exec_result.all = Mock(return_value=[])
cost = get_max_cost_for_model("unknown-model") mock_session.exec.return_value = mock_exec_result
assert cost == 100000 mock_session.get.return_value = None
with patch.object(settings, "fixed_cost_per_request", 100):
with patch.object(settings, "tolerance_percentage", 0):
cost = await get_max_cost_for_model("unknown-model", session=mock_session)
assert cost == 100000
def test_get_max_cost_for_model_disabled() -> None: async def test_get_max_cost_for_model_disabled() -> None:
with patch.object(settings, "fixed_pricing", True): with patch.object(settings, "fixed_pricing", True):
with patch.object(settings, "fixed_cost_per_request", 200): with patch.object(settings, "fixed_cost_per_request", 200):
with patch.object(settings, "tolerance_percentage", 0): with patch.object(settings, "tolerance_percentage", 0):
cost = get_max_cost_for_model("any-model") cost = await get_max_cost_for_model("any-model", session=None)
assert cost == 200000 assert cost == 200000
def test_get_max_cost_for_model_tolerance() -> None: async def test_get_max_cost_for_model_tolerance() -> None:
mock_model = Mock() mock_session = AsyncMock()
mock_model.id = "gpt-4" mock_exec_result = Mock()
mock_model.sats_pricing = Mock() mock_exec_result.all = Mock(return_value=[("gpt-4",)])
mock_model.sats_pricing.max_cost = 500 mock_session.exec.return_value = mock_exec_result
row = Mock()
row.sats_pricing = (
"{" # minimal required fields for Pricing model
'"prompt": 0.0, "completion": 0.0, "request": 0.0, '
'"image": 0.0, "web_search": 0.0, "internal_reasoning": 0.0, '
'"max_cost": 500'
"}"
)
mock_session.get.return_value = row
with patch("routstr.payment.helpers.MODELS", [mock_model]): with patch.object(settings, "fixed_pricing", False):
with patch.object(settings, "fixed_pricing", False): with patch.object(settings, "tolerance_percentage", 10):
with patch.object(settings, "tolerance_percentage", 10): cost = await get_max_cost_for_model("gpt-4", session=mock_session)
cost = get_max_cost_for_model("gpt-4") assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000
assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000