mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fixes
This commit is contained in:
@@ -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
|
||||
|
||||
+21
-44
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+39
-37
@@ -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
|
||||
|
||||
|
||||
|
||||
+1
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user