From 50437a1cc6752b0f23167ab95ed3615de0eff5b4 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Sat, 18 Jul 2026 07:25:52 +0200 Subject: [PATCH] refactor: require explicit settlement identity at the billing seams Make model_obj/provider_fee required (still nullable) on adjust_payment_for_tokens, get_x_cashu_cost and the private pricing helpers so a call site that fails to thread the served candidate is a type error instead of a silent fallback to alias-map re-derivation. calculate_cost keeps its defaults as the one documented fallback seam. Co-Authored-By: Claude Fable 5 --- routstr/auth.py | 4 ++-- routstr/payment/cost_calculation.py | 6 +++--- routstr/upstream/base.py | 6 +++++- .../test_balance_negative_on_cost_overrun.py | 20 +++++++++++++------ tests/integration/test_child_keys.py | 2 +- .../test_free_response_stale_reservation.py | 4 ++-- .../integration/test_reservation_lifecycle.py | 4 +++- tests/unit/test_coverage_base2.py | 4 ++-- 8 files changed, 32 insertions(+), 18 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 3dcd8660..84a57b07 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -774,8 +774,8 @@ async def adjust_payment_for_tokens( response_data: dict, session: AsyncSession, deducted_max_cost: int, - model_obj: "Model | None" = None, - provider_fee: float | None = None, + model_obj: "Model | None", + provider_fee: float | None, ) -> dict: """ Adjusts the payment based on token usage in the response. diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index b9d03d22..37ac15d3 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -323,8 +323,8 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: def _get_pricing_rates( response_data: dict, - model_obj: "Model | None" = None, - provider_fee: float | None = None, + model_obj: "Model | None", + provider_fee: float | None, ) -> tuple[float, float, float, float] | None: """Get configured rates, falling back to LiteLLM's model cost map. @@ -450,7 +450,7 @@ def _calculate_from_usd_cost( cache_creation_tokens: int, output_tokens: int, response_data: dict, - provider_fee: float | None = None, + provider_fee: float | None, ) -> CostData: """Calculate cost from USD figures, deriving input/output split from tokens.""" if provider_fee is None: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 64c03317..f355a1aa 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1707,11 +1707,15 @@ class BaseUpstreamProvider: try: # Finalize with "unknown" model and no usage to release reservation/charge max cost + # (no routed identity here by design: the None usage settles at + # MaxCostData before any pricing lookup can happen). await adjust_payment_for_tokens( key, {"model": "unknown", "usage": None}, session, max_cost, + model_obj=None, + provider_fee=None, ) logger.debug( "Finalized generic streaming payment in background", @@ -3253,7 +3257,7 @@ class BaseUpstreamProvider: self, response_data: dict, max_cost_for_model: int, - model_obj: Model | None = None, + model_obj: Model | None, ) -> MaxCostData | CostData | None: """Calculate cost for X-Cashu payment based on response data. diff --git a/tests/integration/test_balance_negative_on_cost_overrun.py b/tests/integration/test_balance_negative_on_cost_overrun.py index 1cf0f701..e7b8fab8 100644 --- a/tests/integration/test_balance_negative_on_cost_overrun.py +++ b/tests/integration/test_balance_negative_on_cost_overrun.py @@ -82,7 +82,9 @@ async def test_balance_never_negative_when_cost_exceeds_reservation( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + await adjust_payment_for_tokens( + key, response_data, integration_session, deducted_max_cost, None, None + ) await _refresh(integration_session, key) @@ -116,7 +118,9 @@ async def test_balance_floor_at_zero_on_overrun( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + await adjust_payment_for_tokens( + key, response_data, integration_session, deducted_max_cost, None, None + ) await _refresh(integration_session, key) @@ -155,7 +159,9 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + await adjust_payment_for_tokens( + key, response_data, integration_session, deducted_max_cost, None, None + ) await _refresh(integration_session, key) @@ -224,7 +230,7 @@ async def test_concurrent_cost_overruns_never_negative( fresh_key = await session.get(ApiKey, key_hash) assert fresh_key is not None await adjust_payment_for_tokens( - fresh_key, response_data, session, deducted_max_cost + fresh_key, response_data, session, deducted_max_cost, None, None ) # Patch once around the gather: entering the same patch target from @@ -282,7 +288,9 @@ async def test_zero_free_balance_overrun_is_safe( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost) + await adjust_payment_for_tokens( + key, response_data, integration_session, deducted_max_cost, None, None + ) await _refresh(integration_session, key) @@ -348,7 +356,7 @@ async def test_parallel_requests_no_free_inference( fresh_key = await session.get(ApiKey, key_hash) assert fresh_key is not None await adjust_payment_for_tokens( - fresh_key, response_data, session, deducted_max_cost + fresh_key, response_data, session, deducted_max_cost, None, None ) # Patch once around the gather: entering the same patch target from two diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py index 4c141f80..5226d357 100644 --- a/tests/integration/test_child_keys.py +++ b/tests/integration/test_child_keys.py @@ -77,7 +77,7 @@ async def test_child_key_flow(integration_session: AsyncSession) -> None: try: adjustment = await adjust_payment_for_tokens( - child_key_db, response_data, integration_session, 500 + child_key_db, response_data, integration_session, 500, None, None ) assert adjustment["total_msats"] == 400 diff --git a/tests/integration/test_free_response_stale_reservation.py b/tests/integration/test_free_response_stale_reservation.py index 27f52681..00f3c924 100644 --- a/tests/integration/test_free_response_stale_reservation.py +++ b/tests/integration/test_free_response_stale_reservation.py @@ -58,7 +58,7 @@ async def test_overrun_charges_after_reservation_swept( return_value=_cost_data(actual_token_cost), ): await adjust_payment_for_tokens( - key, response_data, integration_session, deducted_max_cost + key, response_data, integration_session, deducted_max_cost, None, None ) await integration_session.refresh(key) @@ -129,7 +129,7 @@ async def test_free_response_path_closed_end_to_end( return_value=_cost_data(actual_token_cost), ): await adjust_payment_for_tokens( - key, response_data, session, deducted_max_cost + key, response_data, session, deducted_max_cost, None, None ) async with create_session() as session: diff --git a/tests/integration/test_reservation_lifecycle.py b/tests/integration/test_reservation_lifecycle.py index ee3fc875..1f60ce9d 100644 --- a/tests/integration/test_reservation_lifecycle.py +++ b/tests/integration/test_reservation_lifecycle.py @@ -120,7 +120,9 @@ async def test_finalise_releases_reservation_and_charges_balance( response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}} with patch("routstr.auth.calculate_cost", return_value=cost_data): - await adjust_payment_for_tokens(key, response_data, integration_session, cost) + await adjust_payment_for_tokens( + key, response_data, integration_session, cost, None, None + ) await integration_session.refresh(key) diff --git a/tests/unit/test_coverage_base2.py b/tests/unit/test_coverage_base2.py index ae7be5ad..1c6a7e46 100644 --- a/tests/unit/test_coverage_base2.py +++ b/tests/unit/test_coverage_base2.py @@ -146,7 +146,7 @@ def test_get_x_cashu_cost_with_usage() -> None: "usage": {"prompt_tokens": 100, "completion_tokens": 50}, } - result = p.get_x_cashu_cost(response_data, 100000) + result = p.get_x_cashu_cost(response_data, 100000, None) # Either returns None (needs more data) or a cost object assert result is not None @@ -157,7 +157,7 @@ def test_get_x_cashu_cost_no_usage() -> None: p = BaseUpstreamProvider("https://api.test.com", "sk-test") response_data = {"model": "gpt-4"} - result = p.get_x_cashu_cost(response_data, 100000) + result = p.get_x_cashu_cost(response_data, 100000, None) # Without usage, uses max_cost assert result is not None