fix other potential reserved balance problems

This commit is contained in:
Shroominic
2026-01-04 23:35:26 +01:00
parent fc8ccf63ba
commit f0c45a7ce4
3 changed files with 155 additions and 23 deletions
+50 -20
View File
@@ -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,
+68 -2
View File
@@ -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,
+37 -1
View File
@@ -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(),