mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
move MODELS to DB
This commit is contained in:
@@ -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
@@ -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,
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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}
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
|
||||||
|
|||||||
Reference in New Issue
Block a user