fix: settle missing-usage at reserved max, add disconnect finalize test

This commit is contained in:
9qeklajc
2026-08-23 12:12:19 +02:00
parent 4aadbd0407
commit 9c7a12808a
3 changed files with 156 additions and 5 deletions
+29 -5
View File
@@ -69,6 +69,30 @@ def _empty_cost(cls: type[CostData] = CostData) -> CostData:
)
def _unmeasured_cost(max_cost: int) -> MaxCostData:
"""Build the bounded fallback for a response whose usage cannot be measured.
Missing usage must NOT settle at zero — that hands out free inference. The
request was authorized up to ``max_cost`` (the reservation), so the safe,
bounded settlement is to charge exactly that. Token components stay zero
because they are genuinely unknown; ``total_msats`` carries the authorized
max so max-cost finalization debits the reservation instead of nothing.
"""
return MaxCostData(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=max(0, max_cost),
total_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
cache_read_msats=0,
cache_creation_msats=0,
)
async def calculate_cost(
response_data: dict,
max_cost: int,
@@ -109,10 +133,10 @@ async def calculate_cost(
if usage is None:
logger.warning(
"No usage data in response — billing at MaxCostData with zero "
"tokens. Dashboard will show this request as `(0+0)`. Most "
"common cause: upstream stream did not include a final usage "
"chunk (OpenAI-compat backends require "
"No usage data in response — settling at the reserved max cost "
"(bounded fallback), not zero. Dashboard will show this request "
"as `(0+0)` tokens. Most common cause: upstream stream did not "
"include a final usage chunk (OpenAI-compat backends require "
"`stream_options.include_usage=true`).",
extra={
"max_cost_msats": max_cost,
@@ -122,7 +146,7 @@ async def calculate_cost(
else None,
},
)
return _empty_cost(MaxCostData)
return _unmeasured_cost(max_cost)
usage_data = response_data.get("usage") or {}
if not isinstance(usage_data, dict):
@@ -91,6 +91,47 @@ async def test_overrun_with_corrupted_aggregate_releases_without_charging(
assert reservation.release_id not in auth._reservation_heartbeats
@pytest.mark.asyncio
async def test_missing_usage_settles_at_reservation_not_zero(
integration_session: AsyncSession,
) -> None:
"""A response with no usable usage data must settle at the reserved max
cost (bounded fallback), never at zero — otherwise the request is free
inference. Exercises the REAL calculate_cost, no patching."""
from routstr.auth import (
adjust_payment_for_tokens,
get_reservation_snapshot,
pay_for_request,
)
reserved = 4_000
key = _make_key(balance=10_000, reserved=0)
key_hash = key.hashed_key
integration_session.add(key)
await integration_session.commit()
await pay_for_request(key, reserved, integration_session)
reservation = await get_reservation_snapshot(key, integration_session)
# No `usage` key at all — the upstream stream dropped its final usage chunk.
response_data = {"model": "test-model"}
result = await adjust_payment_for_tokens(
key,
response_data,
integration_session,
reserved,
reservation_snapshot=reservation,
)
# Charged the authorized max, not zero.
assert result["charged_msats"] == reserved
integration_session.expunge_all()
key_row = await integration_session.get(ApiKey, key_hash)
assert key_row is not None
assert key_row.total_spent == reserved, "missing usage must not be free"
assert key_row.balance == 10_000 - reserved
assert key_row.reserved_balance == 0
@pytest.mark.asyncio
async def test_free_response_path_closed_end_to_end(
integration_session: AsyncSession,
@@ -1,9 +1,11 @@
import asyncio
import json
from collections.abc import AsyncGenerator
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import BackgroundTasks
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import SQLModel
@@ -518,3 +520,87 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() ->
assert second.reserved_balance == 0
await engine.dispose()
@pytest.mark.asyncio
async def test_client_disconnect_midstream_finalizes_and_stops_heartbeat() -> None:
"""A client that aborts the socket mid-stream must not leak its reservation.
Starlette closes the response generator (``aclose``) on disconnect, whose
``finally`` schedules the background finalizer. That finalizer must settle
the reservation (charge the reserved max — usage is unknown), reach a
terminal durable state, and stop the lease heartbeat so the sweeper is not
needed. Driven against a real engine and the real finalizer; the socket
abort is modelled deterministically with ``aclose`` (the exact hook
Starlette invokes) to keep the test CI-stable.
"""
engine = await _engine()
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key", provider_fee=1.0
)
async with AsyncSession(engine, expire_on_commit=False) as session:
key = ApiKey(hashed_key="disconnect-key", balance=1_000)
session.add(key)
await session.commit()
await pay_for_request(key, 500, session)
snapshot = await get_reservation_snapshot(key, session)
assert snapshot.release_id in auth_module._reservation_heartbeats
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
# A live stream that never sends a usage chunk or [DONE]; the client
# disconnects after the first delta.
yield b'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n'
yield b'data: {"choices":[{"delta":{"content":" there"}}]}\n\n'
upstream_response = MagicMock(
status_code=200, headers={"content-type": "text/event-stream"}
)
upstream_response.aiter_bytes = aiter_bytes
background_tasks = BackgroundTasks()
try:
with (
patch(
"routstr.upstream.base.create_session",
side_effect=lambda: AsyncSession(engine, expire_on_commit=False),
),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
auth_module.adjust_payment_for_tokens,
),
):
response = await provider.handle_streaming_chat_completion(
response=upstream_response,
key=key,
max_cost_for_model=500,
background_tasks=background_tasks,
reservation_snapshot=snapshot,
)
iterator = cast(
AsyncGenerator[bytes, None], response.body_iterator
)
await iterator.__anext__() # first chunk reaches the client
await iterator.aclose() # client aborts the socket here
# Starlette runs the response's background tasks after the abort.
for task in background_tasks.tasks:
await task()
finally:
await auth_module._stop_reservation_heartbeat(snapshot.release_id)
async with AsyncSession(engine, expire_on_commit=False) as session:
final_key = await session.get(ApiKey, "disconnect-key")
record = await session.get(ReservationRelease, snapshot.release_id)
assert final_key is not None
# The reservation reached a single terminal outcome; funds are not locked.
assert record is not None and record.status in {"charged", "released"}
assert final_key.reserved_balance == 0
# Unknown usage settles at the reserved max, never free.
assert final_key.total_spent == 500
assert final_key.balance == 500
# The heartbeat is gone — no forever-renewing task on an abandoned request.
assert snapshot.release_id not in auth_module._reservation_heartbeats
await engine.dispose()