revert payment when request failed

This commit is contained in:
9qeklajc
2025-07-25 23:09:16 +02:00
parent bdfb662dc7
commit ce5d655390
2 changed files with 43 additions and 4 deletions
+32 -1
View File
@@ -98,7 +98,7 @@ async def validate_bearer_key(
)
async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> None:
async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int:
cost_per_request = get_max_cost_for_model(model=body["model"])
if key.balance < cost_per_request:
@@ -140,6 +140,37 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> Non
)
await session.refresh(key)
return cost_per_request
async def revert_pay_for_request(
key: ApiKey, session: AsyncSession, cost_per_request: int
) -> None:
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
balance=col(ApiKey.balance) + cost_per_request,
total_spent=col(ApiKey.total_spent) - cost_per_request,
total_requests=col(ApiKey.total_requests) - 1,
)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
raise HTTPException(
status_code=402,
detail={
"error": {
"message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.",
"type": "payment_error",
"code": "payment_error",
}
},
)
await session.refresh(key)
async def adjust_payment_for_tokens(
key: ApiKey, response_data: dict, session: AsyncSession
+11 -3
View File
@@ -15,7 +15,12 @@ from router.payment.helpers import (
)
from router.payment.x_cashu import x_cashu_handler
from .auth import adjust_payment_for_tokens, pay_for_request, validate_bearer_key
from .auth import (
adjust_payment_for_tokens,
pay_for_request,
revert_pay_for_request,
validate_bearer_key,
)
from .cashu import x_cashu_refund
from .db import ApiKey, AsyncSession, create_session, get_session
@@ -305,9 +310,10 @@ async def proxy(
headers = prepare_upstream_headers(dict(request.headers))
return await forward_get_to_upstream(request, path, headers)
cost_per_request = 0
# Only pay for request if we have request body data (for completions endpoints)
if request_body_dict:
await pay_for_request(key, session, request_body_dict)
cost_per_request = await pay_for_request(key, session, request_body_dict)
# Prepare headers for upstream
headers = prepare_upstream_headers(dict(request.headers))
@@ -317,7 +323,9 @@ async def proxy(
request, path, headers, request_body, key, session
)
if response.status_code != 200 and key.refund_address == "X-CASHU":
if response.status_code != 200:
await revert_pay_for_request(key, session, cost_per_request)
refund_token = await x_cashu_refund(key, session, unit)
response = Response(
content=json.dumps(