mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-11 03:28:19 +00:00
revert payment when request failed
This commit is contained in:
+32
-1
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user