From dc0f7d7f3bb98ee0b368cf7a3bf3554591cf4b8a Mon Sep 17 00:00:00 2001 From: Shroominic Date: Wed, 22 Oct 2025 11:54:17 +0800 Subject: [PATCH] fix max_cost discount --- routstr/payment/helpers.py | 25 +++++-------------------- routstr/proxy.py | 2 +- 2 files changed, 6 insertions(+), 21 deletions(-) diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 36beff66..e4893d48 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -9,7 +9,6 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..core import get_logger from ..core.settings import settings from ..wallet import deserialize_token_from_string -from .models import Pricing logger = get_logger(__name__) @@ -152,14 +151,17 @@ async def get_max_cost_for_model( async def calculate_discounted_max_cost( - max_cost_for_model: int, body: dict, session: AsyncSession + max_cost_for_model: int, + body: dict, + model_obj: Any | None = None, ) -> int: """Calculate the discounted max cost for a request using model pricing when available.""" if settings.fixed_pricing: return max_cost_for_model model = body.get("model", "unknown") - model_pricing = await get_model_cost_info(model, session=session) + + model_pricing = model_obj.sats_pricing if model_obj else None if not model_pricing: return max_cost_for_model @@ -215,23 +217,6 @@ def estimate_tokens(messages: list) -> int: return len(str(messages)) // 3 -async def get_model_cost_info(model_id: str, session: AsyncSession) -> Pricing | None: - """Get model pricing info from providers with database overrides.""" - if not model_id or model_id == "unknown": - return None - - from ..proxy import get_upstreams - from ..upstream import get_model_with_override - - upstreams = get_upstreams() - model_obj = await get_model_with_override(model_id, upstreams, session) - - if model_obj and model_obj.sats_pricing: - return model_obj.sats_pricing - - return None - - def create_error_response( error_type: str, message: str, diff --git a/routstr/proxy.py b/routstr/proxy.py index 17e75f6f..86e2cbc9 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -237,7 +237,7 @@ async def proxy( model=model_id, session=session, model_obj=model_obj ) max_cost_for_model = await calculate_discounted_max_cost( - _max_cost_for_model, request_body_dict, session + _max_cost_for_model, request_body_dict, model_obj=model_obj ) check_token_balance(headers, request_body_dict, max_cost_for_model)