diff --git a/migrations/versions/ac10fd366795_add_reservation_releases.py b/migrations/versions/a9bc1d633fa0_add_reservation_release_idempotency_.py similarity index 94% rename from migrations/versions/ac10fd366795_add_reservation_releases.py rename to migrations/versions/a9bc1d633fa0_add_reservation_release_idempotency_.py index 4a3ed0bb..dfec625d 100644 --- a/migrations/versions/ac10fd366795_add_reservation_releases.py +++ b/migrations/versions/a9bc1d633fa0_add_reservation_release_idempotency_.py @@ -1,8 +1,8 @@ """add reservation release idempotency records -Revision ID: ac10fd366795 +Revision ID: a9bc1d633fa0 Revises: d7e8f9a0b1c2 -Create Date: 2026-07-22 22:24:09.482339 +Create Date: 2026-07-24 00:20:27.967658 """ from __future__ import annotations @@ -10,7 +10,7 @@ from __future__ import annotations import sqlalchemy as sa from alembic import op -revision = "ac10fd366795" +revision = "a9bc1d633fa0" down_revision = "d7e8f9a0b1c2" branch_labels = None depends_on = None diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 61eaf078..a8dba7e3 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -796,6 +796,56 @@ class BaseUpstreamProvider: media_type="application/json", ) + async def _release_failed_streaming_reservation( + self, + key: ApiKey, + session: AsyncSession, + reservation_snapshot: ReservationSnapshot | None, + ) -> bool: + """Attempt exact release and suppress unsafe settlement retries.""" + try: + await session.rollback() + snapshot = reservation_snapshot + if snapshot is None: + snapshot = await get_reservation_snapshot(key, session) + released = await release_reservation( + snapshot, + session, + snapshot.reserved_msats, + ) + if not released: + logger.critical( + "Billing reservation could not be released", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "reserved_balance": snapshot.reserved_msats, + }, + ) + # A failed release remains recoverable by the stale-reservation + # sweep. Retrying settlement here could charge after an ambiguous + # database failure or replace the original stream exception. + return True + except asyncio.CancelledError: + # Preserve the exception that triggered billing cleanup. The stream + # propagates it immediately after this helper returns, and stale + # reservation cleanup can recover an interrupted release. + logger.critical( + "Billing reservation release was cancelled", + extra={"key_hash": key.hashed_key[:8] + "..."}, + exc_info=True, + ) + return True + 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, + ) + return True + async def handle_streaming_chat_completion( self, response: httpx.Response, @@ -1043,35 +1093,16 @@ class BaseUpstreamProvider: }, exc_info=True, ) - try: - await session.rollback() - released = await release_reservation( - reservation_snapshot, + # Release is a terminal billing state. Do not enqueue + # finalize_db_only from the generator's finally block + # and charge this request later. + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, session, - max_cost_for_model, - ) - if released: - # Release is a terminal billing state. Do not - # enqueue finalize_db_only from the generator's - # finally block and charge this request later. - usage_finalized = True - else: - 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, + reservation_snapshot, ) + ) raise if usage_chunk_data is None: @@ -1468,23 +1499,23 @@ class BaseUpstreamProvider: reservation_snapshot, ) usage_finalized = True - except Exception as e: - logger.exception( - "Error during Responses API usage finalization", + except BaseException as e: + logger.critical( + "Error during Responses API usage finalization — CRITICAL", extra={ "key_hash": key.hashed_key[:8] + "...", "error": str(e), }, + exc_info=True, ) - cost_data = { - "base_msats": 0, - "input_msats": 0, - "output_msats": 0, - "total_msats": 0, - "total_usd": 0.0, - "input_tokens": 0, - "output_tokens": 0, - } + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + session, + reservation_snapshot, + ) + ) + raise if usage_chunk_data is None: usage_chunk_data = { @@ -1845,9 +1876,23 @@ class BaseUpstreamProvider: ) usage_finalized = True return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() - except Exception: - usage_finalized = True - return None + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise try: async for chunk in response.aiter_bytes(): @@ -2003,8 +2048,23 @@ class BaseUpstreamProvider: usage_finalized = True # Emit the full combined_data as the cost yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode() - except Exception: - pass + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise if not usage_finalized: maybe_cost_event = await finalize_without_usage() @@ -2333,9 +2393,23 @@ class BaseUpstreamProvider: return ( f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" ).encode() - except Exception: - usage_finalized = True - return None + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise try: async for annotated in messages_dispatch.stream_annotated_events( @@ -2410,8 +2484,23 @@ class BaseUpstreamProvider: f"event: cost\ndata: " f"{json.dumps({'cost': cost_data})}\n\n" ).encode() - except Exception: - pass + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + ) + raise if not usage_finalized: cost_event = await finalize_without_usage() @@ -2422,6 +2511,9 @@ class BaseUpstreamProvider: if not usage_finalized: await finalize_without_usage() raise + finally: + if not usage_finalized: + await finalize_without_usage() return StreamingResponse( stream_with_cost(), diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 22ffc7a3..2ae574ab 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import AsyncGenerator from unittest.mock import AsyncMock, MagicMock, patch @@ -9,6 +10,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession import routstr.auth as auth_module from routstr.auth import ( + ReservationSnapshot, adjust_payment_for_tokens, get_reservation_snapshot, pay_for_request, @@ -230,6 +232,7 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge() session_context.__aexit__ = AsyncMock(return_value=None) release = AsyncMock(return_value=True) reservation_snapshot = MagicMock() + reservation_snapshot.reserved_msats = 500 background_tasks = MagicMock() with ( @@ -260,6 +263,162 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge() background_tasks.add_task.assert_not_called() +@pytest.mark.asyncio +@pytest.mark.parametrize( + "release_outcome", + [True, False, RuntimeError("release failed"), asyncio.CancelledError()], +) +async def test_responses_streaming_releases_and_raises_on_billing_failure( + release_outcome: bool | BaseException, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield ( + b'data: {"type":"response.completed","response":{"model":"test",' + b'"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + ) + yield b"data: [DONE]\n\n" + + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = aiter_bytes + key = MagicMock(spec=ApiKey) + key.hashed_key = "responses-key" + 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) + snapshot = ReservationSnapshot( + release_id="responses-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = ( + AsyncMock(side_effect=release_outcome) + if isinstance(release_outcome, BaseException) + else AsyncMock(return_value=release_outcome) + ) + adjust = AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await provider.handle_streaming_responses_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + ) + emitted = bytearray() + with pytest.raises(SQLAlchemyError, match="database unavailable"): + async for chunk in response.body_iterator: + if isinstance(chunk, str): + emitted.extend(chunk.encode()) + else: + emitted.extend(bytes(chunk)) + + assert b'"total_msats": 0' not in emitted + adjust.assert_awaited_once() + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("via_litellm", [False, True]) +@pytest.mark.parametrize( + "release_outcome", + [True, False, RuntimeError("release failed"), asyncio.CancelledError()], +) +async def test_messages_streaming_releases_and_raises_on_billing_failure( + via_litellm: bool, + release_outcome: bool | BaseException, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + key = MagicMock(spec=ApiKey) + key.hashed_key = "messages-key" + 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) + snapshot = ReservationSnapshot( + release_id=f"messages-{'litellm' if via_litellm else 'native'}", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = ( + AsyncMock(side_effect=release_outcome) + if isinstance(release_outcome, BaseException) + else AsyncMock(return_value=release_outcome) + ) + adjust = AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + + async def native_chunks() -> AsyncGenerator[bytes, None]: + yield ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"model":"test","usage":{"input_tokens":1,"output_tokens":0}}}\n\n' + ) + yield b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + + async def litellm_chunks() -> AsyncGenerator[dict, None]: + yield { + "type": "message_start", + "message": { + "model": "test", + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + } + yield {"type": "message_stop"} + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + if via_litellm: + response = provider._stream_litellm_messages( + iterator=litellm_chunks(), + key=key, + max_cost_for_model=500, + requested_model=None, + reservation_snapshot=snapshot, + ) + else: + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = native_chunks + response = await provider.handle_streaming_messages_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + ) + + with pytest.raises(SQLAlchemyError, match="database unavailable"): + async for _ in response.body_iterator: + pass + + adjust.assert_awaited_once() + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + + @pytest.mark.asyncio async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None: engine = await _engine()