This commit is contained in:
Shroominic
2025-09-10 13:28:48 +01:00
parent 73d3613301
commit 4cdbf23ec0
5 changed files with 67 additions and 87 deletions
+1 -1
View File
@@ -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
View File
@@ -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:
+5 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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)