diff --git a/router/auth.py b/router/auth.py index 1cb93dd0..e2feb568 100644 --- a/router/auth.py +++ b/router/auth.py @@ -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 diff --git a/router/proxy.py b/router/proxy.py index 6a5c1528..13f0111e 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -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(