fix: fail safely on streaming billing errors

This commit is contained in:
9qeklajc
2026-07-18 14:42:03 +02:00
parent b3bf1f0e90
commit fa0b366f9a
3 changed files with 134 additions and 15 deletions
+28
View File
@@ -766,6 +766,34 @@ async def revert_pay_for_request(
return True
async def release_reservation(
key: ApiKey,
session: AsyncSession,
reserved_msats: int,
) -> bool:
"""Release a request reservation without charging the key."""
billing_key = await get_billing_key(key, session)
release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= reserved_msats)
.values(reserved_balance=col(ApiKey.reserved_balance) - reserved_msats)
)
result = await session.exec(release_stmt) # type: ignore[call-overload]
if billing_key.hashed_key != key.hashed_key:
child_release_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= reserved_msats)
.values(reserved_balance=col(ApiKey.reserved_balance) - reserved_msats)
)
await session.exec(child_release_stmt) # type: ignore[call-overload]
await session.commit()
return result.rowcount == 1
async def adjust_payment_for_tokens(
key: ApiKey,
response_data: dict,
+29 -15
View File
@@ -13,8 +13,9 @@ import httpx
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from pydantic.v1 import BaseModel
from sqlalchemy.exc import SQLAlchemyError
from ..auth import adjust_payment_for_tokens
from ..auth import adjust_payment_for_tokens, release_reservation
from ..core import get_logger
from ..core.db import (
ApiKey,
@@ -1013,25 +1014,38 @@ class BaseUpstreamProvider:
max_cost_for_model,
)
usage_finalized = True
except Exception as e:
logger.exception(
"Error during usage finalization",
except (HTTPException, SQLAlchemyError) as e:
logger.critical(
"Error during usage finalization — CRITICAL",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
exc_info=True,
)
# Fall back so we still emit a non-zero sats cost downstream.
cost_data = {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
}
try:
await session.rollback()
released = await release_reservation(
fresh_key, session, max_cost_for_model
)
if not released:
logger.critical(
"Billing reservation could not be released",
extra={
"key_hash": key.hashed_key[:8] + "...",
"reserved_balance": fresh_key.reserved_balance,
},
)
except Exception as release_error:
logger.critical(
"Billing reservation release failed",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(release_error),
},
exc_info=True,
)
raise
if usage_chunk_data is None:
if not hasattr(self, "_current_stream_id"):
@@ -0,0 +1,77 @@
from collections.abc import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import release_reservation
from routstr.core.db import ApiKey
from routstr.upstream.base import BaseUpstreamProvider
@pytest.mark.asyncio
async def test_release_reservation_clears_reserved_balance() -> None:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
key = ApiKey(hashed_key="key", balance=1_000, reserved_balance=500)
async with AsyncSession(engine, expire_on_commit=False) as session:
session.add(key)
await session.commit()
assert await release_reservation(key, session, 500) is True
await session.refresh(key)
assert key.reserved_balance == 0
await engine.dispose()
@pytest.mark.asyncio
async def test_streaming_billing_error_releases_reservation_and_propagates() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
yield b"data: [DONE]\n\n"
upstream_response = MagicMock()
upstream_response.status_code = 200
upstream_response.headers = {"content-type": "text/event-stream"}
upstream_response.aiter_bytes = aiter_bytes
key = MagicMock(spec=ApiKey)
key.hashed_key = "test-key-hash"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.rollback = AsyncMock()
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
release = AsyncMock(return_value=True)
with (
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
AsyncMock(side_effect=SQLAlchemyError("database unavailable")),
),
patch("routstr.upstream.base.release_reservation", release),
patch("routstr.upstream.base.create_session", return_value=session_context),
):
response = await provider.handle_streaming_chat_completion(
response=upstream_response,
key=key,
max_cost_for_model=500,
background_tasks=MagicMock(),
)
with pytest.raises(SQLAlchemyError, match="database unavailable"):
async for _ in response.body_iterator:
pass
session.rollback.assert_awaited_once()
release.assert_awaited_once_with(key, session, 500)