From f0c45a7ce46f38f787143abb567f5b239569f1ad Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sun, 4 Jan 2026 23:35:26 +0100 Subject: [PATCH] fix other potential reserved balance problems --- routstr/auth.py | 70 +++++++++++++++++++++++++++----------- routstr/upstream/base.py | 70 ++++++++++++++++++++++++++++++++++++-- routstr/upstream/gemini.py | 38 ++++++++++++++++++++- 3 files changed, 155 insertions(+), 23 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 5b5debdc..b3be04b2 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -441,6 +441,29 @@ async def adjust_payment_for_tokens( }, ) + async def release_reservation_only() -> None: + """Fallback to release reservation without charging when main update fails.""" + try: + release_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost) + ) + await session.exec(release_stmt) # type: ignore[call-overload] + await session.commit() + logger.warning( + "Released reservation without charging (fallback)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, + ) + except Exception as e: + logger.error( + "Failed to release reservation in fallback", + extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."}, + ) + match await calculate_cost(response_data, deducted_max_cost, session): case MaxCostData() as cost: logger.debug( @@ -465,7 +488,7 @@ async def adjust_payment_for_tokens( await session.commit() if result.rowcount == 0: logger.error( - "Failed to finalize max-cost payment - insufficient reserved balance", + "Failed to finalize max-cost payment - retrying reservation release", extra={ "key_hash": key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, @@ -474,6 +497,7 @@ async def adjust_payment_for_tokens( "model": model, }, ) + await release_reservation_only() else: await session.refresh(key) logger.info( @@ -568,13 +592,14 @@ async def adjust_payment_for_tokens( ) else: logger.warning( - "Failed to finalize additional charge (concurrent operation)", + "Failed to finalize additional charge - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", "attempted_charge": total_cost_msats, "model": model, }, ) + await release_reservation_only() else: # Refund some of the base cost refund = abs(cost_difference) @@ -603,7 +628,7 @@ async def adjust_payment_for_tokens( if result.rowcount == 0: logger.error( - "Failed to finalize payment - insufficient reserved balance", + "Failed to finalize payment - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, @@ -612,28 +637,27 @@ async def adjust_payment_for_tokens( "model": model, }, ) - # Still return the cost data even if we couldn't properly finalize - # The reservation was already made, so the user has paid + await release_reservation_only() + else: + cost.total_msats = total_cost_msats + await session.refresh(key) - cost.total_msats = total_cost_msats - await session.refresh(key) - - logger.info( - "Refund processed successfully", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "refunded_amount": refund, - "new_balance": key.balance, - "final_cost": cost.total_msats, - "model": model, - }, - ) + logger.info( + "Refund processed successfully", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "refunded_amount": refund, + "new_balance": key.balance, + "final_cost": cost.total_msats, + "model": model, + }, + ) return cost.dict() case CostDataError() as error: logger.error( - "Cost calculation error during payment adjustment", + "Cost calculation error during payment adjustment - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", "model": model, @@ -641,6 +665,7 @@ async def adjust_payment_for_tokens( "error_code": error.code, }, ) + await release_reservation_only() raise HTTPException( status_code=400, @@ -652,7 +677,12 @@ async def adjust_payment_for_tokens( } }, ) - # Fallback return to satisfy type checker; execution should not reach here + # Fallback: should not reach here, but release reservation just in case + logger.error( + "Unexpected fallback in adjust_payment_for_tokens - releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, + ) + await release_reservation_only() return { "base_msats": deducted_max_cost, "input_msats": 0, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 1ac647e6..5efb69f8 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -447,6 +447,11 @@ class BaseUpstreamProvider: async with create_session() as new_session: fresh_key = await new_session.get(key.__class__, key.hashed_key) if not fresh_key: + logger.warning( + "Key not found when finalizing streaming payment", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = True return None try: fallback: dict = { @@ -475,6 +480,7 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) + usage_finalized = True return None try: @@ -740,6 +746,11 @@ class BaseUpstreamProvider: async with create_session() as new_session: fresh_key = await new_session.get(key.__class__, key.hashed_key) if not fresh_key: + logger.warning( + "Key not found when finalizing Responses API streaming payment", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = True return None try: fallback: dict = { @@ -768,6 +779,7 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) + usage_finalized = True return None try: @@ -892,8 +904,10 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) - await finalize_without_usage() raise + finally: + if not usage_finalized: + await finalize_without_usage() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -1012,6 +1026,44 @@ class BaseUpstreamProvider: ) raise + async def _finalize_generic_streaming_payment( + self, key_hash: str, key_class: type, max_cost: int, path: str + ) -> None: + """Background task to finalize payment for generic streaming requests.""" + async with create_session() as session: + key = await session.get(key_class, key_hash) + if not key: + logger.warning( + "Key not found during background payment finalization", + extra={"key_hash": key_hash[:8] + "..."}, + ) + return + + try: + # Finalize with "unknown" model and no usage to release reservation/charge max cost + await adjust_payment_for_tokens( + key, + {"model": "unknown", "usage": None}, + session, + max_cost, + ) + logger.info( + "Finalized generic streaming payment in background", + extra={ + "path": path, + "key_hash": key_hash[:8] + "...", + }, + ) + except Exception as e: + logger.error( + "Error finalizing generic streaming payment in background", + extra={ + "error": str(e), + "key_hash": key_hash[:8] + "...", + "path": path, + }, + ) + async def forward_request( self, request: Request, @@ -1163,6 +1215,13 @@ class BaseUpstreamProvider: background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) + background_tasks.add_task( + self._finalize_generic_streaming_payment, + key.hashed_key, + key.__class__, + max_cost_for_model, + path, + ) logger.debug( "Streaming non-chat response", @@ -1366,9 +1425,16 @@ class BaseUpstreamProvider: background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) + background_tasks.add_task( + self._finalize_generic_streaming_payment, + key.hashed_key, + key.__class__, + max_cost_for_model, + path, + ) logger.debug( - "Streaming non-chat response", + "Streaming non-Responses API response", extra={ "path": path, "status_code": response.status_code, diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index 79a109c4..a008975a 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -187,11 +187,44 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): ) async def stream_with_cost() -> AsyncGenerator[bytes, None]: + payment_finalized = False + + async def finalize_payment() -> None: + nonlocal payment_finalized + if payment_finalized: + return + from ..auth import adjust_payment_for_tokens + from ..core.db import create_session + + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if fresh_key: + try: + await adjust_payment_for_tokens( + fresh_key, + { + "model": model_obj.id, + "usage": final_usage_data, + }, + new_session, + max_cost_for_model, + ) + payment_finalized = True + except Exception as cost_error: + logger.error( + "Error finalizing Gemini streaming payment in fallback", + extra={ + "error": str(cost_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + try: async for chunk in response_generator: sse_data = f"data: {json.dumps(chunk)}\n\n" yield sse_data.encode() - except Exception as e: logger.error( "Error in Gemini streaming response", @@ -202,6 +235,9 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): }, ) raise + finally: + if not payment_finalized: + await finalize_payment() return StreamingResponse( stream_with_cost(),