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 <noreply@anthropic.com>
This commit is contained in:
Jeroen Ubbink
2026-07-22 15:13:40 +02:00
co-authored by Claude Fable 5
parent 7f5a0cf1ae
commit 50437a1cc6
8 changed files with 32 additions and 18 deletions
+2 -2
View File
@@ -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.
+3 -3
View File
@@ -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:
+5 -1
View File
@@ -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.
@@ -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
+1 -1
View File
@@ -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
@@ -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:
@@ -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)
+2 -2
View File
@@ -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