mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix max_cost discount
This commit is contained in:
@@ -9,7 +9,6 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
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 Pricing
|
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -152,14 +151,17 @@ async def get_max_cost_for_model(
|
|||||||
|
|
||||||
|
|
||||||
async def calculate_discounted_max_cost(
|
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:
|
) -> int:
|
||||||
"""Calculate the discounted max cost for a request using model pricing when available."""
|
"""Calculate the discounted max cost for a request using model pricing when available."""
|
||||||
if settings.fixed_pricing:
|
if settings.fixed_pricing:
|
||||||
return max_cost_for_model
|
return max_cost_for_model
|
||||||
|
|
||||||
model = body.get("model", "unknown")
|
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:
|
if not model_pricing:
|
||||||
return max_cost_for_model
|
return max_cost_for_model
|
||||||
|
|
||||||
@@ -215,23 +217,6 @@ def estimate_tokens(messages: list) -> int:
|
|||||||
return len(str(messages)) // 3
|
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(
|
def create_error_response(
|
||||||
error_type: str,
|
error_type: str,
|
||||||
message: str,
|
message: str,
|
||||||
|
|||||||
+1
-1
@@ -237,7 +237,7 @@ async def proxy(
|
|||||||
model=model_id, session=session, model_obj=model_obj
|
model=model_id, session=session, model_obj=model_obj
|
||||||
)
|
)
|
||||||
max_cost_for_model = await calculate_discounted_max_cost(
|
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)
|
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user