mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-30 15:26:14 +00:00
resolve reviews
This commit is contained in:
+3
-3
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user