From ce5d655390d0f0f8edbd40c4101b60c5eaf666cc Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Fri, 25 Jul 2025 23:03:50 +0200 Subject: [PATCH 1/5] revert payment when request failed --- router/auth.py | 33 ++++++++++++++++++++++++++++++++- router/proxy.py | 14 +++++++++++--- 2 files changed, 43 insertions(+), 4 deletions(-) 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( From c51ce759b70743434dcd9ea449bec9e0fdf29a45 Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Fri, 25 Jul 2025 23:29:29 +0200 Subject: [PATCH 2/5] remove unused code --- router/proxy.py | 25 ------------------------- 1 file changed, 25 deletions(-) diff --git a/router/proxy.py b/router/proxy.py index 13f0111e..813a70d8 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -288,9 +288,6 @@ async def proxy( media_type="application/json", ) - # Check token balance for all requests to get currency unit - unit = check_token_balance(headers, request_body_dict) - # Handle authentication if x_cashu := headers.get("x-cashu", None): return await x_cashu_handler(request, x_cashu, path) @@ -326,28 +323,6 @@ async def proxy( 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( - { - "error": { - "message": "Error forwarding request to upstream", - "type": "upstream_error", - "code": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=response.status_code, - media_type="application/json", - ) - response.headers["X-Cashu"] = refund_token - return response - - if key.refund_address == "X-CASHU": - refund_token = await x_cashu_refund(key, session, unit) - response.headers["X-Cashu"] = refund_token - return response From 98b5fee6eb0513171f11fe398672aec62d7e3d9c Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Fri, 25 Jul 2025 23:30:46 +0200 Subject: [PATCH 3/5] clean up --- router/proxy.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/router/proxy.py b/router/proxy.py index 813a70d8..4963d458 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -9,7 +9,6 @@ from fastapi.responses import Response, StreamingResponse from router.payment.helpers import ( UPSTREAM_BASE_URL, - check_token_balance, create_error_response, prepare_upstream_headers, ) @@ -21,7 +20,6 @@ from .auth import ( revert_pay_for_request, validate_bearer_key, ) -from .cashu import x_cashu_refund from .db import ApiKey, AsyncSession, create_session, get_session proxy_router = APIRouter() From 9af2e9646586ef7b5ae8965b9ef9b9359c2bf7b5 Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Sun, 27 Jul 2025 01:03:59 +0200 Subject: [PATCH 4/5] revert changes --- router/proxy.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/router/proxy.py b/router/proxy.py index 4963d458..ea1d6eb6 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -9,6 +9,7 @@ from fastapi.responses import Response, StreamingResponse from router.payment.helpers import ( UPSTREAM_BASE_URL, + check_token_balance, create_error_response, prepare_upstream_headers, ) @@ -285,6 +286,8 @@ async def proxy( status_code=400, media_type="application/json", ) + # Check token balance for all requests to get currency unit + _ = check_token_balance(headers, request_body_dict) # Handle authentication if x_cashu := headers.get("x-cashu", None): From c893ce41b6b2aaefcde8d47d11b6def03281daa3 Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Sun, 27 Jul 2025 10:56:46 +0200 Subject: [PATCH 5/5] fmt --- router/proxy.py | 1 + 1 file changed, 1 insertion(+) diff --git a/router/proxy.py b/router/proxy.py index ea1d6eb6..9c2848fe 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -286,6 +286,7 @@ async def proxy( status_code=400, media_type="application/json", ) + # Check token balance for all requests to get currency unit _ = check_token_balance(headers, request_body_dict)