mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge pull request #170 from Routstr/token-discount-max-cost
max cost discount
This commit is contained in:
@@ -50,6 +50,7 @@ class Settings(BaseSettings):
|
|||||||
fixed_per_1k_output_tokens: int = Field(default=0, env="FIXED_PER_1K_OUTPUT_TOKENS")
|
fixed_per_1k_output_tokens: int = Field(default=0, env="FIXED_PER_1K_OUTPUT_TOKENS")
|
||||||
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")
|
||||||
|
|
||||||
# Network
|
# Network
|
||||||
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
|
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import math
|
||||||
from typing import Mapping
|
from typing import Mapping
|
||||||
|
|
||||||
from fastapi import HTTPException, Response
|
from fastapi import HTTPException, Response
|
||||||
@@ -7,7 +8,7 @@ from fastapi.requests import Request
|
|||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
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
|
from .models import MODELS, Pricing
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -80,7 +81,7 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_max_cost_for_model(model: str, tolerance_percentage: int = 1) -> int:
|
def get_max_cost_for_model(model: str) -> 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",
|
||||||
@@ -115,7 +116,11 @@ def get_max_cost_for_model(model: str, tolerance_percentage: int = 1) -> int:
|
|||||||
|
|
||||||
for m in MODELS:
|
for m in MODELS:
|
||||||
if m.id == model:
|
if m.id == model:
|
||||||
max_cost = m.sats_pricing.max_cost * 1000 * (1 - tolerance_percentage / 100) # type: ignore
|
max_cost = (
|
||||||
|
m.sats_pricing.max_cost # type: ignore
|
||||||
|
* 1000
|
||||||
|
* (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},
|
||||||
@@ -132,6 +137,83 @@ def get_max_cost_for_model(model: str, tolerance_percentage: int = 1) -> int:
|
|||||||
return settings.fixed_cost_per_request * 1000
|
return settings.fixed_cost_per_request * 1000
|
||||||
|
|
||||||
|
|
||||||
|
def calculate_discounted_max_cost(max_cost_for_model: int, body: dict) -> int:
|
||||||
|
"""Calculate the discounted max cost for a request."""
|
||||||
|
original_max_cost_msats = max_cost_for_model
|
||||||
|
model = body.get("model", "unknown")
|
||||||
|
|
||||||
|
if settings.fixed_pricing:
|
||||||
|
return max_cost_for_model
|
||||||
|
|
||||||
|
if not (model_pricing := get_model_cost_info(model)):
|
||||||
|
return max_cost_for_model
|
||||||
|
|
||||||
|
tol = settings.tolerance_percentage
|
||||||
|
tol_factor = max(0.0, 1 - float(tol) / 100.0)
|
||||||
|
max_prompt_allowed_sats = model_pricing.max_prompt_cost * tol_factor
|
||||||
|
max_completion_allowed_sats = model_pricing.max_completion_cost * tol_factor
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"Discount estimation context",
|
||||||
|
extra={
|
||||||
|
"model": model,
|
||||||
|
"tolerance_pct": tol,
|
||||||
|
"tol_factor": tol_factor,
|
||||||
|
"start_max_cost_msats": original_max_cost_msats,
|
||||||
|
"model_max_cost_sats": model_pricing.max_cost,
|
||||||
|
"model_max_prompt_cost_sats": model_pricing.max_prompt_cost,
|
||||||
|
"model_max_completion_cost_sats": model_pricing.max_completion_cost,
|
||||||
|
"input_rate_sats_per_token": model_pricing.prompt,
|
||||||
|
"output_rate_sats_per_token": model_pricing.completion,
|
||||||
|
"max_prompt_allowed_sats": max_prompt_allowed_sats,
|
||||||
|
"max_completion_allowed_sats": max_completion_allowed_sats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
if messages := body.get("messages"):
|
||||||
|
prompt_tokens = estimate_tokens(messages)
|
||||||
|
estimated_prompt_delta_sats = (
|
||||||
|
max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt
|
||||||
|
)
|
||||||
|
if estimated_prompt_delta_sats >= 0:
|
||||||
|
max_cost_for_model = max_cost_for_model - math.floor(
|
||||||
|
estimated_prompt_delta_sats * 1000
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
max_cost_for_model = max_cost_for_model + math.ceil(
|
||||||
|
-estimated_prompt_delta_sats * 1000
|
||||||
|
)
|
||||||
|
|
||||||
|
if max_tokens := body.get("max_tokens"):
|
||||||
|
estimated_completion_delta_sats = (
|
||||||
|
max_completion_allowed_sats - max_tokens * model_pricing.completion
|
||||||
|
)
|
||||||
|
if estimated_completion_delta_sats >= 0:
|
||||||
|
max_cost_for_model = max_cost_for_model - math.floor(
|
||||||
|
estimated_completion_delta_sats * 1000
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
max_cost_for_model = max_cost_for_model + math.ceil(
|
||||||
|
-estimated_completion_delta_sats * 1000
|
||||||
|
)
|
||||||
|
|
||||||
|
print("max_cost_for_model", max_cost_for_model)
|
||||||
|
|
||||||
|
return max(0, max_cost_for_model)
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_tokens(messages: list) -> int:
|
||||||
|
return len(str(messages)) // 3
|
||||||
|
|
||||||
|
|
||||||
|
def get_model_cost_info(model_id: str) -> Pricing | None:
|
||||||
|
if not model_id or model_id == "unknown":
|
||||||
|
return None
|
||||||
|
|
||||||
|
model = next((m for m in MODELS if m.id == model_id), None)
|
||||||
|
return model.sats_pricing if model else None # type: ignore
|
||||||
|
|
||||||
|
|
||||||
def create_error_response(
|
def create_error_response(
|
||||||
error_type: str,
|
error_type: str,
|
||||||
message: str,
|
message: str,
|
||||||
|
|||||||
+46
-14
@@ -30,6 +30,8 @@ class Pricing(BaseModel):
|
|||||||
image: float
|
image: float
|
||||||
web_search: float
|
web_search: float
|
||||||
internal_reasoning: float
|
internal_reasoning: float
|
||||||
|
max_prompt_cost: float = 0.0 # in sats not msats
|
||||||
|
max_completion_cost: float = 0.0 # in sats not msats
|
||||||
max_cost: float = 0.0 # in sats not msats
|
max_cost: float = 0.0 # in sats not msats
|
||||||
|
|
||||||
|
|
||||||
@@ -111,7 +113,7 @@ def load_models() -> list[Model]:
|
|||||||
try:
|
try:
|
||||||
with models_path.open("r") as f:
|
with models_path.open("r") as f:
|
||||||
data = json.load(f)
|
data = json.load(f)
|
||||||
return [Model(**model) for model in data.get("models", [])]
|
return [Model(**model) for model in data.get("models", [])] # type: ignore
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error loading models from {models_path}: {e}")
|
logger.error(f"Error loading models from {models_path}: {e}")
|
||||||
# Fall through to auto-generation
|
# Fall through to auto-generation
|
||||||
@@ -130,7 +132,7 @@ def load_models() -> list[Model]:
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
logger.info(f"Successfully fetched {len(models_data)} models from OpenRouter API")
|
logger.info(f"Successfully fetched {len(models_data)} models from OpenRouter API")
|
||||||
return [Model(**model) for model in models_data]
|
return [Model(**model) for model in models_data] # type: ignore
|
||||||
|
|
||||||
|
|
||||||
MODELS = load_models()
|
MODELS = load_models()
|
||||||
@@ -143,26 +145,54 @@ async def update_sats_pricing() -> None:
|
|||||||
for model in MODELS:
|
for model in MODELS:
|
||||||
model.sats_pricing = Pricing(
|
model.sats_pricing = Pricing(
|
||||||
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
|
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
|
||||||
)
|
) # type: ignore
|
||||||
mspp = model.sats_pricing.prompt
|
mspp = model.sats_pricing.prompt
|
||||||
mspc = model.sats_pricing.completion
|
mspc = model.sats_pricing.completion
|
||||||
if (tp := model.top_provider) and (
|
if (tp := model.top_provider) and (
|
||||||
tp.context_length or tp.max_completion_tokens
|
tp.context_length or tp.max_completion_tokens
|
||||||
):
|
):
|
||||||
if (cl := model.top_provider.context_length) and (
|
if (cl := tp.context_length) and (mct := tp.max_completion_tokens):
|
||||||
mct := model.top_provider.max_completion_tokens
|
max_prompt_cost = (cl - mct) * mspp
|
||||||
):
|
max_completion_cost = mct * mspc
|
||||||
model.sats_pricing.max_cost = (cl - mct) * mspp + mct * mspc
|
model.sats_pricing.max_prompt_cost = max_prompt_cost
|
||||||
elif cl := model.top_provider.context_length:
|
model.sats_pricing.max_completion_cost = max_completion_cost
|
||||||
model.sats_pricing.max_cost = cl * 0.8 * mspp + cl * 0.2 * mspc
|
model.sats_pricing.max_cost = (
|
||||||
elif mct := model.top_provider.max_completion_tokens:
|
max_prompt_cost + max_completion_cost
|
||||||
model.sats_pricing.max_cost = mct * 4 * mspp + mct * mspc
|
)
|
||||||
|
elif cl := tp.context_length:
|
||||||
|
max_prompt_cost = cl * 0.8 * mspp
|
||||||
|
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:
|
||||||
|
max_prompt_cost = mct * 4 * mspp
|
||||||
|
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
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
model.sats_pricing.max_cost = 1_000_000 * mspp + 32_000 * mspc
|
max_prompt_cost = 1_000_000 * mspp
|
||||||
|
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:
|
elif model.context_length:
|
||||||
model.sats_pricing.max_cost = (
|
max_prompt_cost = (
|
||||||
model.sats_pricing.prompt * model.context_length * 0.8
|
model.sats_pricing.prompt * model.context_length * 0.8
|
||||||
) + (model.sats_pricing.completion * model.context_length * 0.2)
|
)
|
||||||
|
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:
|
else:
|
||||||
p = model.sats_pricing.prompt * 1_000_000
|
p = model.sats_pricing.prompt * 1_000_000
|
||||||
c = model.sats_pricing.completion * 32_000
|
c = model.sats_pricing.completion * 32_000
|
||||||
@@ -170,6 +200,8 @@ async def update_sats_pricing() -> None:
|
|||||||
i = model.sats_pricing.image * 100
|
i = model.sats_pricing.image * 100
|
||||||
w = model.sats_pricing.web_search * 1000
|
w = model.sats_pricing.web_search * 1000
|
||||||
ir = model.sats_pricing.internal_reasoning * 100
|
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
|
model.sats_pricing.max_cost = p + c + r + i + w + ir
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
break
|
||||||
|
|||||||
+5
-1
@@ -17,6 +17,7 @@ from .core import get_logger
|
|||||||
from .core.db import ApiKey, AsyncSession, create_session, get_session
|
from .core.db import ApiKey, AsyncSession, create_session, get_session
|
||||||
from .core.settings import settings
|
from .core.settings import settings
|
||||||
from .payment.helpers import (
|
from .payment.helpers import (
|
||||||
|
calculate_discounted_max_cost,
|
||||||
check_token_balance,
|
check_token_balance,
|
||||||
create_error_response,
|
create_error_response,
|
||||||
get_max_cost_for_model,
|
get_max_cost_for_model,
|
||||||
@@ -554,7 +555,10 @@ 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 = get_max_cost_for_model(model=model)
|
||||||
|
max_cost_for_model = calculate_discounted_max_cost(
|
||||||
|
_max_cost_for_model, request_body_dict
|
||||||
|
)
|
||||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||||
|
|
||||||
# Handle authentication
|
# Handle authentication
|
||||||
|
|||||||
@@ -509,6 +509,10 @@ async def integration_app(
|
|||||||
mint_url = os.environ.get("CASHU_MINTS", "http://localhost:3338")
|
mint_url = os.environ.get("CASHU_MINTS", "http://localhost:3338")
|
||||||
from routstr.core.settings import settings as _settings
|
from routstr.core.settings import settings as _settings
|
||||||
|
|
||||||
|
# Passthrough discounted max cost to avoid dependence on MODELS in tests
|
||||||
|
def _passthrough_discount(max_cost_for_model: int, body: dict) -> int:
|
||||||
|
return max_cost_for_model
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch("routstr.core.db.engine", integration_engine),
|
patch("routstr.core.db.engine", integration_engine),
|
||||||
patch.object(_settings, "cashu_mints", [mint_url]),
|
patch.object(_settings, "cashu_mints", [mint_url]),
|
||||||
@@ -522,6 +526,10 @@ async def integration_app(
|
|||||||
patch("websockets.connect") as mock_websockets,
|
patch("websockets.connect") as mock_websockets,
|
||||||
patch("routstr.payment.price.btc_usd_ask_price", return_value=50000.0),
|
patch("routstr.payment.price.btc_usd_ask_price", return_value=50000.0),
|
||||||
patch("routstr.payment.price.sats_usd_ask_price", return_value=0.0005),
|
patch("routstr.payment.price.sats_usd_ask_price", return_value=0.0005),
|
||||||
|
patch(
|
||||||
|
"routstr.payment.helpers.calculate_discounted_max_cost",
|
||||||
|
side_effect=_passthrough_discount,
|
||||||
|
),
|
||||||
):
|
):
|
||||||
# Configure the WebSocket mock for discovery service - fast failure for performance tests
|
# Configure the WebSocket mock for discovery service - fast failure for performance tests
|
||||||
async def mock_websocket_connect(*args: Any, **kwargs: Any) -> None:
|
async def mock_websocket_connect(*args: Any, **kwargs: Any) -> None:
|
||||||
|
|||||||
@@ -17,22 +17,25 @@ def test_get_max_cost_for_model_known() -> None:
|
|||||||
|
|
||||||
with patch("routstr.payment.helpers.MODELS", [mock_model]):
|
with patch("routstr.payment.helpers.MODELS", [mock_model]):
|
||||||
with patch.object(settings, "fixed_pricing", False):
|
with patch.object(settings, "fixed_pricing", False):
|
||||||
cost = get_max_cost_for_model("gpt-4", tolerance_percentage=0)
|
with patch.object(settings, "tolerance_percentage", 0):
|
||||||
assert cost == 500000 # 500 sats * 1000 = msats
|
cost = get_max_cost_for_model("gpt-4")
|
||||||
|
assert cost == 500000 # 500 sats * 1000 = msats
|
||||||
|
|
||||||
|
|
||||||
def test_get_max_cost_for_model_unknown() -> None:
|
def test_get_max_cost_for_model_unknown() -> None:
|
||||||
with patch("routstr.payment.helpers.MODELS", []):
|
with patch("routstr.payment.helpers.MODELS", []):
|
||||||
with patch.object(settings, "fixed_cost_per_request", 100):
|
with patch.object(settings, "fixed_cost_per_request", 100):
|
||||||
cost = get_max_cost_for_model("unknown-model", tolerance_percentage=0)
|
with patch.object(settings, "tolerance_percentage", 0):
|
||||||
assert cost == 100000
|
cost = get_max_cost_for_model("unknown-model")
|
||||||
|
assert cost == 100000
|
||||||
|
|
||||||
|
|
||||||
def test_get_max_cost_for_model_disabled() -> None:
|
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):
|
||||||
cost = get_max_cost_for_model("any-model", tolerance_percentage=0)
|
with patch.object(settings, "tolerance_percentage", 0):
|
||||||
assert cost == 200000
|
cost = get_max_cost_for_model("any-model")
|
||||||
|
assert cost == 200000
|
||||||
|
|
||||||
|
|
||||||
def test_get_max_cost_for_model_tolerance() -> None:
|
def test_get_max_cost_for_model_tolerance() -> None:
|
||||||
@@ -43,5 +46,6 @@ def test_get_max_cost_for_model_tolerance() -> None:
|
|||||||
|
|
||||||
with patch("routstr.payment.helpers.MODELS", [mock_model]):
|
with patch("routstr.payment.helpers.MODELS", [mock_model]):
|
||||||
with patch.object(settings, "fixed_pricing", False):
|
with patch.object(settings, "fixed_pricing", False):
|
||||||
cost = get_max_cost_for_model("gpt-4", tolerance_percentage=10)
|
with patch.object(settings, "tolerance_percentage", 10):
|
||||||
assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000
|
cost = get_max_cost_for_model("gpt-4")
|
||||||
|
assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000
|
||||||
|
|||||||
Reference in New Issue
Block a user