mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-09 02:54:37 +00:00
fix: fail safely on streaming billing errors
This commit is contained in:
@@ -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
@@ -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)
|
||||
Reference in New Issue
Block a user