From fa0b366f9a38212070e1dc85df6cf6755cb80e0c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 18 Jul 2026 14:42:03 +0200 Subject: [PATCH] fix: fail safely on streaming billing errors --- routstr/auth.py | 28 +++++++ routstr/upstream/base.py | 44 +++++++---- .../test_streaming_billing_finalization.py | 77 +++++++++++++++++++ 3 files changed, 134 insertions(+), 15 deletions(-) create mode 100644 tests/unit/test_streaming_billing_finalization.py diff --git a/routstr/auth.py b/routstr/auth.py index 23ad6920..1655eea2 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 306aca14..9b2ecdfa 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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"): diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py new file mode 100644 index 00000000..d599b268 --- /dev/null +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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)