diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 7dbab415..86495728 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -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): diff --git a/tests/integration/test_free_response_stale_reservation.py b/tests/integration/test_free_response_stale_reservation.py index 80f0e294..77b4579b 100644 --- a/tests/integration/test_free_response_stale_reservation.py +++ b/tests/integration/test_free_response_stale_reservation.py @@ -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, diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 51ab2ccf..94ef65a4 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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()