diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index e7cfdb87..a4ded317 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -270,11 +270,11 @@ async def calculate_cost( input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates # Truthiness is not the question: `NaN` and a negative rate are both truthy - # and sailed past this gate into the token math. + # and sailed past this gate into the token math, while a rate of zero is a + # price — free — and reading it as a missing one charged the whole + # reservation for a request the model serves for nothing. rates = (input_rate, output_rate, cache_read_rate, cache_creation_rate) - if not all(is_usable_rate(rate) for rate in rates) or not ( - input_rate and output_rate - ): + if not all(is_usable_rate(rate) for rate in rates): logger.warning( "No usable token pricing — billing at flat MaxCostData. " "Token counts %s in the upstream response but cannot be " diff --git a/tests/unit/test_pricing_rate_validation.py b/tests/unit/test_pricing_rate_validation.py index fab4e6d5..f0390c9b 100644 --- a/tests/unit/test_pricing_rate_validation.py +++ b/tests/unit/test_pricing_rate_validation.py @@ -92,6 +92,30 @@ async def test_unusable_token_rate_falls_back_to_max_cost(bad_rate: float) -> No assert cost.total_msats == 1234 +@pytest.mark.parametrize( + ("prompt", "completion", "expected_msats"), + [(0.0, 0.0, 0), (0.0, 2e-06, 1), (1e-06, 0.0, 1)], + ids=["free", "free-input", "free-output"], +) +@pytest.mark.asyncio +async def test_a_rate_of_zero_is_billed_as_free_not_as_missing( + prompt: float, completion: float, expected_msats: int +) -> None: + """Zero is a price, and the request must be billed on it. + + The gate that decides a model has no token pricing was a truthiness test, so + a free rate read as an absent one and the request was charged the whole + reservation instead — on a model priced at zero for that side, which is a + price the catalog serves and the router routes. + """ + model = _model(Pricing(prompt=prompt, completion=completion)) + + cost = await calculate_cost(_usage_response(), max_cost=1234, model_obj=model) + + assert isinstance(cost, CostData) + assert cost.total_msats == expected_msats + + @pytest.mark.parametrize( "junk", [float("inf"), float("nan"), "Infinity"], ids=["inf", "nan", "inf-string"] )