From c104345561420f64a2f81f7d21ccd403c91716fc Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 19:55:22 +0200 Subject: [PATCH] fix: treat non-finite token counts as zero instead of raising in billing --- routstr/payment/usage.py | 17 +++-- tests/unit/test_non_finite_token_counts.py | 82 ++++++++++++++++++++++ 2 files changed, 95 insertions(+), 4 deletions(-) create mode 100644 tests/unit/test_non_finite_token_counts.py diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 08d675ad..d69ab55c 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -38,6 +38,8 @@ names do not collide, so a single union parser is safe; a vendor whose fields would genuinely conflict needs a dedicated branch here. """ +import math + from pydantic.v1 import BaseModel @@ -51,18 +53,25 @@ class NormalizedUsage(BaseModel): def parse_token_count(value: object) -> int: - """Parse a token count from various formats (int, float, str, bool).""" + """Parse a token count from various formats (int, float, str, bool). + + ``json.loads`` accepts bare ``Infinity``/``NaN`` and overflows ``1e999`` to + ``inf``, so an upstream can put them on the wire. ``int()`` raises on both, + which would turn a billing path into a 500; reject them like + ``is_usable_rate`` does instead. + """ if isinstance(value, bool): return 0 if isinstance(value, int): return max(0, value) if isinstance(value, float): - return max(0, int(value)) + return max(0, int(value)) if math.isfinite(value) else 0 if isinstance(value, str): try: - return max(0, int(float(value))) - except ValueError: + parsed = float(value) + except (ValueError, OverflowError): return 0 + return max(0, int(parsed)) if math.isfinite(parsed) else 0 return 0 diff --git a/tests/unit/test_non_finite_token_counts.py b/tests/unit/test_non_finite_token_counts.py new file mode 100644 index 00000000..572dfe87 --- /dev/null +++ b/tests/unit/test_non_finite_token_counts.py @@ -0,0 +1,82 @@ +"""Non-finite token counts must not crash the billing path. + +``json.loads`` accepts bare ``Infinity``/``NaN`` and overflows ``1e999`` to +``inf``, so an upstream can put them on the wire. ``int()`` raises on both. +""" + +import json +import os +from typing import Any +from unittest.mock import patch + +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") +os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") + +from routstr.core.settings import settings +from routstr.payment.cost_calculation import CostData, calculate_cost +from routstr.payment.usage import parse_token_count + +NON_FINITE = [ + float("inf"), + float("-inf"), + float("nan"), + 1e999, + "Infinity", + "NaN", + "-Infinity", + "1e999", +] + + +@pytest.fixture(autouse=True) +def _fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> Any: + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001) + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5): + yield + + +@pytest.mark.parametrize("value", NON_FINITE) +def test_parse_token_count_rejects_non_finite(value: Any) -> None: + assert parse_token_count(value) == 0 + + +def test_parse_token_count_still_parses_ordinary_values() -> None: + assert parse_token_count(42) == 42 + assert parse_token_count("42") == 42 + assert parse_token_count(42.9) == 42 + assert parse_token_count("42.9") == 42 + assert parse_token_count(True) == 0 + assert parse_token_count(-5) == 0 + assert parse_token_count("not a number") == 0 + assert parse_token_count(None) == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "raw", + [ + '{"usage": {"prompt_tokens": Infinity, "completion_tokens": 10}}', + '{"usage": {"prompt_tokens": 1e999, "completion_tokens": 10}}', + '{"usage": {"prompt_tokens": 1000, "completion_tokens": NaN}}', + '{"usage": {"prompt_tokens": "Infinity", "completion_tokens": "10"}}', + ], +) +async def test_billing_settles_a_response_with_non_finite_usage(raw: str) -> None: + """The billing entry point parses the decoded wire body without raising + and bills only the finite component.""" + response = json.loads(raw) + + result = await calculate_cost(response, max_cost=100_000) + + assert isinstance(result, CostData) + usage = response["usage"] + finite_input = parse_token_count(usage["prompt_tokens"]) + finite_output = parse_token_count(usage["completion_tokens"]) + assert result.input_tokens == finite_input + assert result.output_tokens == finite_output + assert 0 <= result.total_msats <= 100_000