resolve reviews

This commit is contained in:
9qeklajc
2026-07-24 00:24:24 +02:00
parent dbe7a53afd
commit 2ed20b1b85
3 changed files with 303 additions and 52 deletions
@@ -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
+141 -49
View File
@@ -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(),
@@ -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()