From 4cdbf23ec06026323982b109ea758703eab75747 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Wed, 10 Sep 2025 13:28:48 +0100 Subject: [PATCH] fixes --- routstr/core/main.py | 2 +- routstr/payment/helpers.py | 65 +++++++++++--------------------- routstr/payment/models.py | 9 +++-- routstr/payment/x_cashu.py | 76 +++++++++++++++++++------------------- routstr/proxy.py | 2 +- 5 files changed, 67 insertions(+), 87 deletions(-) diff --git a/routstr/core/main.py b/routstr/core/main.py index b3906eef..eb1515a9 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -13,8 +13,8 @@ from ..nip91 import announce_provider from ..payment.models import ( ensure_models_bootstrapped, models_router, - update_sats_pricing, refresh_models_periodically, + update_sats_pricing, ) from ..proxy import proxy_router from ..wallet import periodic_payout diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index c7d304f8..892aec59 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -154,25 +154,15 @@ async def get_max_cost_for_model( return max(settings.min_request_msat, settings.fixed_cost_per_request * 1000) -def calculate_discounted_max_cost( +async 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.""" - original_max_cost_msats = max_cost_for_model - model = body.get("model", "unknown") - - if settings.fixed_pricing: + """Calculate the discounted max cost for a request using model pricing when available.""" + if settings.fixed_pricing or session is None: return max_cost_for_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 + model = body.get("model", "unknown") + model_pricing = await get_model_cost_info(model, session=session) if not model_pricing: return max_cost_for_model @@ -181,22 +171,7 @@ def calculate_discounted_max_cost( 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, - }, - ) + adjusted = max_cost_for_model if messages := body.get("messages"): prompt_tokens = estimate_tokens(messages) @@ -204,28 +179,30 @@ def calculate_discounted_max_cost( 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 - ) + adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000) else: - max_cost_for_model = max_cost_for_model + math.ceil( - -estimated_prompt_delta_sats * 1000 - ) + adjusted = adjusted + 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 - ) + adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000) else: - max_cost_for_model = max_cost_for_model + math.ceil( - -estimated_completion_delta_sats * 1000 - ) + adjusted = adjusted + math.ceil(-estimated_completion_delta_sats * 1000) - return max(0, max_cost_for_model) + logger.debug( + "Discounted max cost computed", + extra={ + "model": model, + "original_msats": max_cost_for_model, + "adjusted_msats": adjusted, + "tolerance_pct": tol, + }, + ) + + return max(0, adjusted) def estimate_tokens(messages: list) -> int: diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 8624cb2c..f21da6bf 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -1,5 +1,6 @@ import asyncio import json +import random from pathlib import Path from urllib.request import urlopen @@ -372,8 +373,8 @@ async def update_sats_pricing() -> None: logger.error(f"Error updating sats pricing: {e}") try: 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)) + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) except asyncio.CancelledError: break @@ -437,8 +438,8 @@ async def refresh_models_periodically() -> None: 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)) + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) except asyncio.CancelledError: break diff --git a/routstr/payment/x_cashu.py b/routstr/payment/x_cashu.py index 0ec614b2..f671fbd8 100644 --- a/routstr/payment/x_cashu.py +++ b/routstr/payment/x_cashu.py @@ -7,6 +7,7 @@ from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from ..core import get_logger +from ..core.db import create_session from ..core.settings import settings from ..wallet import recieve_token, send_token from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost @@ -553,43 +554,44 @@ async def get_cost( extra={"model": model, "has_usage": "usage" in response_data}, ) - match await calculate_cost(response_data, max_cost_for_model, None): - case MaxCostData() as cost: - logger.debug( - "Using max cost pricing", - extra={"model": model, "max_cost_msats": cost.total_msats}, - ) - return cost - case CostData() as cost: - logger.debug( - "Using token-based pricing", - extra={ - "model": model, - "total_cost_msats": cost.total_msats, - "input_msats": cost.input_msats, - "output_msats": cost.output_msats, - }, - ) - return cost - case CostDataError() as error: - logger.error( - "Cost calculation error", - extra={ - "model": model, - "error_message": error.message, - "error_code": error.code, - }, - ) - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": error.message, - "type": "invalid_request_error", - "code": error.code, - } - }, - ) + async with create_session() as session: + match await calculate_cost(response_data, max_cost_for_model, session): + case MaxCostData() as cost: + logger.debug( + "Using max cost pricing", + extra={"model": model, "max_cost_msats": cost.total_msats}, + ) + return cost + case CostData() as cost: + logger.debug( + "Using token-based pricing", + extra={ + "model": model, + "total_cost_msats": cost.total_msats, + "input_msats": cost.input_msats, + "output_msats": cost.output_msats, + }, + ) + return cost + case CostDataError() as error: + logger.error( + "Cost calculation error", + extra={ + "model": model, + "error_message": error.message, + "error_code": error.code, + }, + ) + raise HTTPException( + status_code=400, + detail={ + "error": { + "message": error.message, + "type": "invalid_request_error", + "code": error.code, + } + }, + ) return None diff --git a/routstr/proxy.py b/routstr/proxy.py index c4573581..aebf80a1 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -556,7 +556,7 @@ async def proxy( model = request_body_dict.get("model", "unknown") _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 = await calculate_discounted_max_cost( _max_cost_for_model, request_body_dict, session ) check_token_balance(headers, request_body_dict, max_cost_for_model)