fix max_cost race contition when price changes

This commit is contained in:
Shroominic
2025-08-22 17:22:48 -03:00
parent 6c53c0661c
commit d7e35887de
4 changed files with 44 additions and 58 deletions
+3 -6
View File
@@ -272,10 +272,10 @@ async def validate_bearer_key(
)
async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int:
async def pay_for_request(
key: ApiKey, cost_per_request: int, session: AsyncSession
) -> int:
"""Process payment for a request."""
model = body["model"]
cost_per_request = get_max_cost_for_model(model=model)
logger.info(
"Processing payment for request",
@@ -283,7 +283,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int
"key_hash": key.hashed_key[:8] + "...",
"current_balance": key.balance,
"required_cost": cost_per_request,
"model": model,
"sufficient_balance": key.balance >= cost_per_request,
},
)
@@ -297,7 +296,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int
"reserved_balance": key.reserved_balance,
"required": cost_per_request,
"shortfall": cost_per_request - key.total_balance,
"model": model,
},
)
@@ -366,7 +364,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int
"new_balance": key.balance,
"total_spent": key.total_spent,
"total_requests": key.total_requests,
"model": model,
},
)
-24
View File
@@ -19,30 +19,6 @@ if not UPSTREAM_BASE_URL:
raise ValueError("Please set the UPSTREAM_BASE_URL environment variable")
def get_cost_per_request(model: str | None = None) -> int:
"""Get the cost per request for a given model."""
logger.debug(
"Calculating cost per request",
extra={
"model": model,
"model_based_pricing": MODEL_BASED_PRICING,
"has_models": bool(MODELS),
},
)
if MODEL_BASED_PRICING and MODELS and model:
cost = get_max_cost_for_model(model=model)
logger.debug(
"Using model-based cost", extra={"model": model, "cost_msats": cost}
)
return cost
logger.debug(
"Using default cost per request", extra={"cost_msats": COST_PER_REQUEST}
)
return COST_PER_REQUEST
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
if x_cashu := headers.get("x-cashu", None):
cashu_token = x_cashu
+36 -22
View File
@@ -9,18 +9,13 @@ from fastapi.responses import Response, StreamingResponse
from ..core import get_logger
from ..wallet import recieve_token, send_token
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
from .helpers import (
UPSTREAM_BASE_URL,
create_error_response,
get_max_cost_for_model,
prepare_upstream_headers,
)
from .helpers import UPSTREAM_BASE_URL, create_error_response, prepare_upstream_headers
logger = get_logger(__name__)
async def x_cashu_handler(
request: Request, x_cashu_token: str, path: str
request: Request, x_cashu_token: str, path: str, max_cost_for_model: int
) -> Response | StreamingResponse:
"""Handle X-Cashu token payment requests."""
logger.info(
@@ -44,7 +39,9 @@ async def x_cashu_handler(
extra={"amount": amount, "unit": unit, "path": path, "mint": mint},
)
return await forward_to_upstream(request, path, headers, amount, unit)
return await forward_to_upstream(
request, path, headers, amount, unit, max_cost_for_model
)
except Exception as e:
error_message = str(e)
logger.error(
@@ -96,7 +93,12 @@ async def x_cashu_handler(
async def forward_to_upstream(
request: Request, path: str, headers: dict, amount: int, unit: str
request: Request,
path: str,
headers: dict,
amount: int,
unit: str,
max_cost_for_model: int,
) -> Response | StreamingResponse:
"""Forward request to upstream and handle the response."""
if path.startswith("v1/"):
@@ -188,7 +190,9 @@ async def forward_to_upstream(
extra={"path": path, "amount": amount, "unit": unit},
)
result = await handle_x_cashu_chat_completion(response, amount, unit)
result = await handle_x_cashu_chat_completion(
response, amount, unit, max_cost_for_model
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
result.background = background_tasks
@@ -232,7 +236,7 @@ async def forward_to_upstream(
async def handle_x_cashu_chat_completion(
response: httpx.Response, amount: int, unit: str
response: httpx.Response, amount: int, unit: str, max_cost_for_model: int
) -> StreamingResponse | Response:
"""Handle both streaming and non-streaming chat completion responses with token-based pricing."""
logger.debug(
@@ -256,10 +260,12 @@ async def handle_x_cashu_chat_completion(
)
if is_streaming:
return await handle_streaming_response(content_str, response, amount, unit)
return await handle_streaming_response(
content_str, response, amount, unit, max_cost_for_model
)
else:
return await handle_non_streaming_response(
content_str, response, amount, unit
content_str, response, amount, unit, max_cost_for_model
)
except Exception as e:
@@ -281,7 +287,11 @@ async def handle_x_cashu_chat_completion(
async def handle_streaming_response(
content_str: str, response: httpx.Response, amount: int, unit: str
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
) -> StreamingResponse:
"""Handle Server-Sent Events (SSE) streaming response."""
logger.debug(
@@ -335,7 +345,7 @@ async def handle_streaming_response(
response_data = {"usage": usage_data, "model": model}
try:
cost_data = await get_cost(response_data)
cost_data = await get_cost(response_data, max_cost_for_model)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
@@ -403,7 +413,11 @@ async def handle_streaming_response(
async def handle_non_streaming_response(
content_str: str, response: httpx.Response, amount: int, unit: str
content_str: str,
response: httpx.Response,
amount: int,
unit: str,
max_cost_for_model: int,
) -> Response:
"""Handle regular JSON response."""
logger.debug(
@@ -414,7 +428,7 @@ async def handle_non_streaming_response(
try:
response_json = json.loads(content_str)
cost_data = await get_cost(response_json)
cost_data = await get_cost(response_json, max_cost_for_model)
if not cost_data:
logger.error(
@@ -520,21 +534,21 @@ async def handle_non_streaming_response(
)
async def get_cost(response_data: dict) -> MaxCostData | CostData | None:
async def get_cost(
response_data: dict, max_cost_for_model: int
) -> MaxCostData | CostData | None:
"""
Adjusts the payment based on token usage in the response.
This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response.
"""
model = response_data.get("model", "unknown")
model = response_data.get("model", None)
logger.debug(
"Calculating cost for response",
extra={"model": model, "has_usage": "usage" in response_data},
)
max_cost = get_max_cost_for_model(model=model)
match calculate_cost(response_data, max_cost):
match calculate_cost(response_data, max_cost_for_model):
case MaxCostData() as cost:
logger.debug(
"Using max cost pricing",
+5 -6
View File
@@ -19,7 +19,7 @@ from .payment.helpers import (
UPSTREAM_BASE_URL,
check_token_balance,
create_error_response,
get_cost_per_request,
get_max_cost_for_model,
prepare_upstream_headers,
)
from .payment.x_cashu import x_cashu_handler
@@ -501,9 +501,8 @@ async def proxy(
media_type="application/json",
)
max_cost_for_model = get_cost_per_request(
model=request_body_dict.get("model", None)
)
model = request_body_dict.get("model", "unknown")
max_cost_for_model = get_max_cost_for_model(model=model)
check_token_balance(headers, request_body_dict, max_cost_for_model)
# Handle authentication
@@ -515,7 +514,7 @@ async def proxy(
"token_preview": x_cashu[:20] + "..." if len(x_cashu) > 20 else x_cashu,
},
)
return await x_cashu_handler(request, x_cashu, path)
return await x_cashu_handler(request, x_cashu, path, max_cost_for_model)
elif auth := headers.get("authorization", None):
logger.debug(
@@ -557,7 +556,7 @@ async def proxy(
)
try:
await pay_for_request(key, session, request_body_dict)
await pay_for_request(key, max_cost_for_model, session)
logger.info(
"Payment processed successfully",
extra={