From a7d4e2832df84faea6ea1011171cdd30a1dce866 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 24 Aug 2026 01:55:15 +0200 Subject: [PATCH] fix: gate upstream-reported cost behind a per-provider trust policy MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A positive cost reported by an upstream pre-empted token pricing with no check on who reported it. Because the bearer overrun path settles min(chargeable, total) against the key's balance minus sibling reservations, any configured or chained provider that controls an accepted cost field could bill far beyond the reservation it was authorized against and drain the key. The mirror case bled the operator: an under-reporting provider with omitted token usage settled at its own low number. Reported cost is now honoured only for provider types that opt in via BaseUpstreamProvider.trusts_reported_cost, which defaults to off. A new provider therefore prices from tokens until someone deliberately approves it. PPQ.AI (BYOK — only PPQ knows the user's upstream bill) and OpenRouter (per-request sub-provider routing) are approved; chained Routstr peers and generic/custom rows are not. This is a policy gate, not a clamp: PPQ BYOK legitimately settles above the reservation and still does. The flag is threaded from the serving provider instance through adjust_payment_for_tokens to calculate_cost alongside provider_fee, at every streaming and non-streaming settlement site in the upstream base provider, so the two paths agree. Two supporting changes: - _coerce_usd rejects non-finite input. Infinity previously survived the clamp and reached math.ceil, where the OverflowError was swallowed by the USD path's broad handler — correct by accident. NaN was already folded to zero by max() comparison semantics; it is now explicit. Nothing about int/float msats arithmetic changed, so billing amounts are unaffected. - A trusted provider's cost is compared against what its own reported tokens would have been priced at, and against the reservation. Ratios outside the bounds are logged in both directions. They are not clamped: the legitimate BYOK spread is wide enough that clamping would mis-bill real traffic, so the goal is that a mis-report is visible rather than silent. Existing USD-path tests now declare their provider as cost-reporting; no assertion was changed. --- routstr/auth.py | 12 +- routstr/payment/cost_calculation.py | 151 ++++++++-- routstr/upstream/base.py | 20 ++ routstr/upstream/openrouter.py | 4 + routstr/upstream/ppqai.py | 4 + tests/integration/test_failover_billing.py | 18 +- tests/unit/test_cost_calculation_caching.py | 52 ++-- tests/unit/test_messages_litellm_dispatch.py | 68 ++--- tests/unit/test_reported_cost_trust.py | 277 +++++++++++++++++++ tests/unit/test_settlement_identity.py | 34 +-- 10 files changed, 528 insertions(+), 112 deletions(-) create mode 100644 tests/unit/test_reported_cost_trust.py diff --git a/routstr/auth.py b/routstr/auth.py index 0dfa38e4..00ac6d87 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -1298,6 +1298,8 @@ async def adjust_payment_for_tokens( model_obj: "Model | None" = None, provider_fee: float | None = None, reservation_snapshot: ReservationSnapshot | None = None, + *, + trusts_reported_cost: bool = False, ) -> dict: """ Adjusts the payment based on token usage in the response. @@ -1308,6 +1310,10 @@ async def adjust_payment_for_tokens( through to ``calculate_cost`` so billing uses the serving candidate's pricing instead of re-deriving it from the response's model string. + ``trusts_reported_cost`` carries the serving provider type's cost-reporting + policy; it defaults to off so a caller that cannot name its provider never + lets the upstream price itself. + The response's usage object is normalized with the default union parser in ``calculate_cost``. """ @@ -1375,7 +1381,11 @@ async def adjust_payment_for_tokens( ) calculated_cost = await calculate_cost( - response_data, deducted_max_cost, model_obj, provider_fee + response_data, + deducted_max_cost, + model_obj, + provider_fee, + trusts_reported_cost=trusts_reported_cost, ) if not isinstance(calculated_cost, CostDataError): if not await _claim_reservation_for_charge(reservation, session): diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 86495728..43472edd 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -21,6 +21,13 @@ __all__ = [ logger = get_logger(__name__) +# Bounds for surfacing a reported cost that does not resemble what the tokens +# would have been priced at. These only alert — PPQ.AI BYOK legitimately +# settles at roughly ten times the reservation, so clamping here would +# under-charge real traffic. +REPORTED_COST_TOKEN_RATIO_BOUNDS = (0.2, 5.0) +REPORTED_COST_RESERVATION_RATIO = 20.0 + class CostData(BaseModel): base_msats: int @@ -98,6 +105,8 @@ async def calculate_cost( max_cost: int, model_obj: "Model | None" = None, provider_fee: float | None = None, + *, + trusts_reported_cost: bool = False, ) -> CostData | MaxCostData | CostDataError: """Calculate the cost of an API request based on token usage. @@ -113,6 +122,12 @@ async def calculate_cost( pricing already carries the fee baked in). Without it, the fee is re-derived from the response's model string, which yields the best-ranked provider's fee. + trusts_reported_cost: Whether the serving provider type is approved to + name its own price (``BaseUpstreamProvider.trusts_reported_cost``). + A reported cost pre-empts token pricing and the overrun path will + settle it against the key's whole unreserved balance, so an + unapproved provider's cost fields are discarded and the request is + priced from the tokens it actually reported. Returns: Cost data or error information @@ -159,6 +174,18 @@ async def calculate_cost( # Try USD cost first usd_cost = _resolve_usd_cost(usage_data, response_data) + if usd_cost > 0 and not trusts_reported_cost: + logger.warning( + "Upstream reported a cost but its provider type is not approved " + "to price its own requests — discarding the reported cost and " + "billing from token usage.", + extra={ + "model": response_data.get("model", "unknown"), + "reported_usd_cost": usd_cost, + "max_cost_msats": max_cost, + }, + ) + usd_cost = 0.0 if usd_cost > 0: truly_empty = ( input_tokens == 0 @@ -209,29 +236,10 @@ async def calculate_cost( cost_details.get("output_cost") or cost_details.get("upstream_inference_completions_cost") ) - cache_pricing_rates: tuple[float, float, float, float] | None = None - if cache_read_tokens > 0 or cache_creation_tokens > 0: - try: - cache_pricing_rates = _get_pricing_rates( - response_data, model_obj, provider_fee - ) - except ValueError: - logger.warning( - "Cache pricing unavailable for USD cost breakdown; " - "leaving cache cost components unknown", - extra={"model": response_data.get("model", "unknown")}, - ) - if cache_pricing_rates is None and settings.fixed_pricing: - fixed_input_rate = ( - float(settings.fixed_per_1k_input_tokens) * 1000.0 - ) - cache_pricing_rates = ( - fixed_input_rate, - float(settings.fixed_per_1k_output_tokens) * 1000.0, - fixed_input_rate, - fixed_input_rate, - ) - return _calculate_from_usd_cost( + cache_pricing_rates = _usd_path_pricing_rates( + response_data, model_obj, provider_fee + ) + cost = _calculate_from_usd_cost( usd_cost, input_usd, output_usd, @@ -243,6 +251,10 @@ async def calculate_cost( provider_fee, cache_pricing_rates, ) + _flag_reported_cost_anomalies( + cost, max_cost, cache_pricing_rates, response_data + ) + return cost except Exception as e: logger.warning( "Error calculating cost from usage data", @@ -319,9 +331,15 @@ def _coerce_usd(value: object) -> float: if not isinstance(value, (int, float, str)): return 0.0 try: - return max(0.0, float(value)) + parsed = float(value) except (TypeError, ValueError): return 0.0 + # Infinity survives the clamp below and then overflows ``math.ceil``; + # NaN is folded to zero here rather than relying on ``max`` comparison + # semantics to do it. + if not math.isfinite(parsed): + return 0.0 + return max(0.0, parsed) def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: @@ -366,6 +384,91 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: return 0.0 +def _usd_path_pricing_rates( + response_data: dict, + model_obj: "Model | None", + provider_fee: float | None, +) -> tuple[float, float, float, float] | None: + """Best-effort token rates for the USD path's cache split and plausibility check. + + Unlike the token-priced path this must never fail the request: the upstream + total stays authoritative whether or not local rates are known. + """ + rates: tuple[float, float, float, float] | None = None + try: + rates = _get_pricing_rates(response_data, model_obj, provider_fee) + except ValueError: + logger.warning( + "Local pricing unavailable for USD cost breakdown; " + "leaving cache cost components unknown", + extra={"model": response_data.get("model", "unknown")}, + ) + if rates is None and settings.fixed_pricing: + fixed_input_rate = float(settings.fixed_per_1k_input_tokens) * 1000.0 + return ( + fixed_input_rate, + float(settings.fixed_per_1k_output_tokens) * 1000.0, + fixed_input_rate, + fixed_input_rate, + ) + return rates + + +def _flag_reported_cost_anomalies( + cost: CostData, + max_cost: int, + pricing_rates: tuple[float, float, float, float] | None, + response_data: dict, +) -> None: + """Alert on a trusted provider's cost that does not match its own tokens. + + Both directions matter: over-reporting drains the client up to its + authorization, under-reporting bleeds the operator. Neither is clamped — + a trusted provider's total is still what gets billed — because the + legitimate BYOK spread is wide enough that a clamp would mis-bill real + traffic. This exists so the mis-report is visible rather than silent. + """ + model = response_data.get("model", "unknown") + + if pricing_rates is not None: + input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates + token_priced_msats = ( + cost.input_tokens / 1000 * input_rate + + cost.output_tokens / 1000 * output_rate + + cost.cache_read_input_tokens / 1000 * cache_read_rate + + cost.cache_creation_input_tokens / 1000 * cache_creation_rate + ) + low, high = REPORTED_COST_TOKEN_RATIO_BOUNDS + if token_priced_msats > 0: + ratio = cost.total_msats / token_priced_msats + if ratio < low or ratio > high: + logger.warning( + "Upstream-reported cost is implausible against token " + "pricing for the same response — billing it as reported, " + "but the provider is over- or under-reporting.", + extra={ + "model": model, + "reported_msats": cost.total_msats, + "token_priced_msats": token_priced_msats, + "ratio": ratio, + "input_tokens": cost.input_tokens, + "output_tokens": cost.output_tokens, + }, + ) + + if max_cost > 0 and cost.total_msats > max_cost * REPORTED_COST_RESERVATION_RATIO: + logger.warning( + "Upstream-reported cost is far above the reservation it was " + "authorized against — the overrun will settle against the key's " + "unreserved balance.", + extra={ + "model": model, + "reported_msats": cost.total_msats, + "max_cost_msats": max_cost, + }, + ) + + def _get_pricing_rates( response_data: dict, model_obj: "Model | None", diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 4080be0f..640e1e3b 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -240,6 +240,12 @@ class BaseUpstreamProvider: platform_url: str | None = None supports_anthropic_messages: bool = False + # Whether this provider type may price its own requests. A reported cost + # pre-empts token pricing and settles against the key's unreserved + # balance, so it is opt-in per provider type: only upstreams whose billing + # we cannot reconstruct locally (BYOK, dynamic sub-provider routing) get + # it. Chained and operator-configured peers stay on token pricing. + trusts_reported_cost: bool = False # When None, the prefix is detected from `base_url` at dispatch time # (see `get_litellm_provider_prefix`). Subclasses set this to lock the # provider regardless of URL. @@ -1047,6 +1053,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True except Exception: @@ -1220,6 +1227,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True except BaseException as e: @@ -1362,6 +1370,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) await session.refresh(key) @@ -1497,6 +1506,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True except Exception: @@ -1627,6 +1637,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True except BaseException as e: @@ -1773,6 +1784,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) await session.refresh(key) @@ -1885,6 +1897,7 @@ class BaseUpstreamProvider: model_obj=model_obj, provider_fee=provider_fee, reservation_snapshot=reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) logger.debug( "Finalized generic streaming payment in background", @@ -1977,6 +1990,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() @@ -2143,6 +2157,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) self.inject_cost_metadata( @@ -2234,6 +2249,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2346,6 +2362,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2501,6 +2518,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) usage_finalized = True return ( @@ -2588,6 +2606,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + trusts_reported_cost=self.trusts_reported_cost, ) self.inject_cost_metadata( combined_data, cost_data, fresh_key @@ -3566,6 +3585,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + trusts_reported_cost=self.trusts_reported_cost, ): case MaxCostData() as cost: logger.debug( diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1caeaa5c..fe0b4f7f 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -17,6 +17,10 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): platform_url = "https://openrouter.ai/settings/keys" supports_anthropic_messages = True litellm_provider_prefix = "openrouter/" + # OpenRouter picks the serving sub-provider per request, so its reported + # cost is the only accurate price — our model rates describe the router's + # advertised range, not what this particular request was billed at. + trusts_reported_cost = True def _apply_provider_field(self, response_json: object) -> None: """Stamp the ``provider`` field for OpenRouter responses. diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 50ab4532..29503ae8 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -44,6 +44,10 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): # provider-attested usage extractor/model binding for it. Keep EHBP disabled # until a ConfidentialInferenceProfile can bill it without max-cost fallback. supports_ehbp = False + # Under BYOK the inference is billed against the user's own upstream key; + # only PPQ knows what that cost was, so its report is authoritative and + # legitimately exceeds the reservation. + trusts_reported_cost = True def __init__(self, api_key: str, provider_fee: float = 1.0): super().__init__( diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index a0c31401..4a260447 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -59,9 +59,17 @@ def _make_model( class _StaticProvider(BaseUpstreamProvider): """Upstream provider with a fixed model catalog and no remote refresh.""" - def __init__(self, base_url: str, api_key: str, fee: float, model: Model) -> None: + def __init__( + self, + base_url: str, + api_key: str, + fee: float, + model: Model, + trusts_reported_cost: bool = False, + ) -> None: super().__init__(base_url, api_key, fee) self.provider_type = "custom" + self.trusts_reported_cost = trusts_reported_cost self._static_model = model def get_cached_models(self) -> list[Model]: @@ -369,18 +377,24 @@ async def test_version_suffixed_model_id_routes( async def fee_split_provider_maps( patched_db_engine: None, ) -> AsyncGenerator[None, None]: - """Same-tail providers whose fees differ; the serving one charges 1.5x.""" + """Same-tail providers whose fees differ; the serving one charges 1.5x. + + Both are approved to report their own cost, which is what puts billing on + the USD path at all. + """ cheap = _StaticProvider( CHEAP_BASE_URL, "key-cheap", 1.0, _make_model("dual-model", 0.001, 0.002), + trusts_reported_cost=True, ) expensive = _StaticProvider( EXPENSIVE_BASE_URL, "key-expensive", 1.5, _make_model("dual-model", 0.005, 0.010), + trusts_reported_cost=True, ) async for _ in _install_providers([cheap, expensive]): yield diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 65ad5091..3703929b 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -46,8 +46,8 @@ async def test_openai_cache_subtraction() -> None: "completion_tokens": 100, "prompt_tokens_details": { "cached_tokens": 1000 # ← Extracted separately - } - } + }, + }, } result = await calculate_cost(response, max_cost=100000) @@ -70,7 +70,7 @@ async def test_anthropic_cache_additive(mock_fixed_pricing: None) -> None: "output_tokens": 100, "cache_creation_input_tokens": 1500, # ← Additive, not included above "cache_read_input_tokens": 0, - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -94,8 +94,8 @@ async def test_cache_read_exceeds_prompt_tokens(mock_fixed_pricing: None) -> Non "completion_tokens": 50, "prompt_tokens_details": { "cached_tokens": 150 # ← Invalid! Greater than prompt - } - } + }, + }, } result = await calculate_cost(response, max_cost=100000) @@ -120,8 +120,8 @@ async def test_malformed_cache_tokens_coerce_to_zero(mock_fixed_pricing: None) - "cache_read_input_tokens": "-50", # ← String, negative "prompt_tokens_details": { "cached_tokens": "invalid" # ← Non-numeric string - } - } + }, + }, } result = await calculate_cost(response, max_cost=100000) @@ -143,7 +143,7 @@ async def test_anthropic_cache_not_subtracted(mock_fixed_pricing: None) -> None: "input_tokens": 500, "completion_tokens": 100, "cache_read_input_tokens": 200, # ← Additive, don't subtract - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -164,10 +164,8 @@ async def test_only_cache_read_tokens(mock_fixed_pricing: None) -> None: "usage": { "prompt_tokens": 0, "completion_tokens": 50, - "prompt_tokens_details": { - "cached_tokens": 1000 - } - } + "prompt_tokens_details": {"cached_tokens": 1000}, + }, } result = await calculate_cost(response, max_cost=100000) @@ -190,7 +188,7 @@ async def test_only_cache_creation_tokens(mock_fixed_pricing: None) -> None: "output_tokens": 100, "cache_creation_input_tokens": 2000, "cache_read_input_tokens": 0, - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -214,7 +212,7 @@ async def test_both_cache_read_and_creation(mock_fixed_pricing: None) -> None: "output_tokens": 100, "cache_creation_input_tokens": 2000, "cache_read_input_tokens": 500, - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -237,7 +235,7 @@ async def test_token_field_fallback_order(mock_fixed_pricing: None) -> None: "usage": { "input_tokens": 250, "completion_tokens": 50, - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -258,7 +256,7 @@ async def test_float_token_values_coerced_to_int(mock_fixed_pricing: None) -> No "prompt_tokens": 100.7, # Float "completion_tokens": 50.3, # Float "prompt_tokens_details": {"cached_tokens": 25.9}, # Float - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -281,7 +279,7 @@ async def test_boolean_cache_tokens_coerced_to_zero(mock_fixed_pricing: None) -> "prompt_tokens": 100, "completion_tokens": 50, "cache_read_input_tokens": True, # Boolean - } + }, } result = await calculate_cost(response, max_cost=100000) @@ -301,10 +299,8 @@ async def test_zero_cache_tokens(mock_fixed_pricing: None) -> None: "usage": { "prompt_tokens": 100, "completion_tokens": 50, - "prompt_tokens_details": { - "cached_tokens": 0 - } - } + "prompt_tokens_details": {"cached_tokens": 0}, + }, } result = await calculate_cost(response, max_cost=100000) @@ -427,7 +423,7 @@ async def test_truly_empty_usd_cost_response_is_refunded( "total_cost": 0.01, # non-zero USD cost despite no tokens }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) assert result.total_msats == 0 # full refund @@ -455,7 +451,7 @@ async def test_cache_read_only_usd_cost_response_is_billed( "total_cost": 0.01, # non-zero USD cost }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) # NOT refunded — the USD cost is billed in full. Pinning the exact value @@ -494,7 +490,7 @@ async def test_small_usd_cost_components_sum_to_rounded_total( }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) assert result.total_msats == expected_msats @@ -525,7 +521,7 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None: }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) assert result.input_msats == 995 @@ -590,6 +586,7 @@ async def test_usd_cache_breakdown_matches_token_priced_path( max_cost=100_000, model_obj=model, provider_fee=1.0, + trusts_reported_cost=True, ) assert isinstance(token_result, CostData) @@ -650,6 +647,7 @@ async def test_usd_cache_breakdown_does_not_absorb_total_rounding_remainder( max_cost=100_000, model_obj=model, provider_fee=1.0, + trusts_reported_cost=True, ) assert isinstance(token_result, CostData) @@ -687,7 +685,7 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None: }, }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) # The fix bills upstream_inference_cost + byok_fee (~0.047 USD → ~940k @@ -719,7 +717,7 @@ async def test_ppq_byok_fee_only_would_undercharge() -> None: "prompt_tokens_details": {"cached_tokens": 159301}, }, } - result = await calculate_cost(response, max_cost=100000) + result = await calculate_cost(response, max_cost=100000, trusts_reported_cost=True) assert isinstance(result, CostData) # Without upstream_inference_cost, only the fee is billed — the old bug. diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index ec93d555..06c7c50f 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -48,7 +48,9 @@ def _make_model( return Model( id=model_id, name=model_id, - forwarded_model_id=forwarded_model_id if forwarded_model_id is not None else model_id, + forwarded_model_id=forwarded_model_id + if forwarded_model_id is not None + else model_id, created=0, description="", context_length=8192, @@ -173,9 +175,7 @@ def test_events_from_chunk_handles_bytes_chunks() -> None: def test_events_from_chunk_handles_str_chunks() -> None: provider = _make_provider() - events, buf = provider._events_from_chunk( - 'event: a\ndata: {"type":"a"}\n\n', b"" - ) + events, buf = provider._events_from_chunk('event: a\ndata: {"type":"a"}\n\n', b"") assert events == [{"type": "a"}] assert buf == b"" @@ -218,9 +218,7 @@ def test_base_provider_resolves_prefix_from_base_url() -> None: ) assert groq.get_litellm_provider_prefix() == "groq/" - xai = BaseUpstreamProvider( - base_url="https://api.x.ai/v1", api_key="sk-test" - ) + xai = BaseUpstreamProvider(base_url="https://api.x.ai/v1", api_key="sk-test") assert xai.get_litellm_provider_prefix() == "xai/" deepseek = BaseUpstreamProvider( @@ -228,9 +226,7 @@ def test_base_provider_resolves_prefix_from_base_url() -> None: ) assert deepseek.get_litellm_provider_prefix() == "deepseek/" - unknown = BaseUpstreamProvider( - base_url="https://example.com/v1", api_key="sk-test" - ) + unknown = BaseUpstreamProvider(base_url="https://example.com/v1", api_key="sk-test") assert unknown.get_litellm_provider_prefix() == "openai/" @@ -581,6 +577,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: model_obj: Any = None, provider_fee: Any = None, reservation_snapshot: Any = None, + trusts_reported_cost: bool = False, ) -> dict: captured_cost_call["combined_data"] = combined_data captured_cost_call["max_cost"] = max_cost @@ -685,6 +682,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: model_obj: Any = None, provider_fee: Any = None, reservation_snapshot: Any = None, + trusts_reported_cost: bool = False, ) -> dict: captured["combined_data"] = combined_data captured["reservation_snapshot"] = reservation_snapshot @@ -1254,7 +1252,7 @@ async def test_aggregator_parses_tool_use_input_json_delta() -> None: { "type": "content_block_delta", "index": 0, - "delta": {"type": "input_json_delta", "partial_json": ' 7}'}, + "delta": {"type": "input_json_delta", "partial_json": " 7}"}, }, {"type": "content_block_stop", "index": 0}, { @@ -1282,18 +1280,18 @@ async def test_aggregator_parses_sse_byte_chunks() -> None: provider = _make_provider() sse = ( - b'event: message_start\n' + b"event: message_start\n" b'data: {"type":"message_start","message":{"id":"m1","type":"message",' b'"role":"assistant","model":"x","content":[],"usage":{"input_tokens":1,"output_tokens":0}}}\n\n' - b'event: content_block_start\n' + b"event: content_block_start\n" b'data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n' - b'event: content_block_delta\n' + b"event: content_block_delta\n" b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}\n\n' - b'event: content_block_stop\n' + b"event: content_block_stop\n" b'data: {"type":"content_block_stop","index":0}\n\n' - b'event: message_delta\n' + b"event: message_delta\n" b'data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":2}}\n\n' - b'event: message_stop\n' + b"event: message_stop\n" b'data: {"type":"message_stop"}\n\n' ) @@ -1362,11 +1360,13 @@ async def test_dispatch_always_streams_upstream_and_aggregates_for_non_streaming "litellm.anthropic.messages.acreate", new=AsyncMock(side_effect=fake_acreate), ): - client_stream, result, requested_model = ( - await provider._dispatch_anthropic_messages( - request_body=body, - model_obj=model, - ) + ( + client_stream, + result, + requested_model, + ) = await provider._dispatch_anthropic_messages( + request_body=body, + model_obj=model, ) # Upstream was streamed regardless of client preference. @@ -1404,13 +1404,9 @@ async def test_dispatch_uses_url_detected_prefix_for_fireworks_custom_row() -> N "litellm.anthropic.messages.acreate", new=AsyncMock(side_effect=fake_acreate), ): - await provider._dispatch_anthropic_messages( - request_body=body, model_obj=model - ) + await provider._dispatch_anthropic_messages(request_body=body, model_obj=model) - assert captured_kwargs["model"] == ( - "fireworks_ai/accounts/fireworks/models/glm-5" - ) + assert captured_kwargs["model"] == ("fireworks_ai/accounts/fireworks/models/glm-5") assert captured_kwargs["api_base"] == "https://api.fireworks.ai/inference/v1" @@ -1441,9 +1437,7 @@ async def test_x_cashu_mint_unreachable_returns_503( model = _make_model() request = _make_request() - with patch( - "routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error) - ): + with patch("routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)): handler = getattr(provider, handler_name) response = await handler( request=request, @@ -1539,9 +1533,7 @@ async def test_x_cashu_error_code_is_stable_string( model = _make_model() request = _make_request() - with patch( - "routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error) - ): + with patch("routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)): handler = getattr(provider, handler_name) response = await handler( request=request, @@ -1661,9 +1653,7 @@ async def test_x_cashu_echoes_token_only_when_recoverable( model = _make_model() request = _make_request() - with patch( - "routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error) - ): + with patch("routstr.upstream.base.recieve_token", new=AsyncMock(side_effect=error)): handler = getattr(provider, handler_name) response = await handler( request=request, @@ -1707,7 +1697,9 @@ async def test_x_cashu_zero_value_rejected_not_forwarded( patch.object( provider, forward_attr, - new=AsyncMock(side_effect=AssertionError("must not forward a zero-value token")), + new=AsyncMock( + side_effect=AssertionError("must not forward a zero-value token") + ), ), ): handler = getattr(provider, handler_name) diff --git a/tests/unit/test_reported_cost_trust.py b/tests/unit/test_reported_cost_trust.py new file mode 100644 index 00000000..73df21e8 --- /dev/null +++ b/tests/unit/test_reported_cost_trust.py @@ -0,0 +1,277 @@ +"""Tests for the per-provider trust policy on upstream-reported cost. + +An upstream that names its own price pre-empts token pricing entirely, and the +bearer overrun path will then spend whatever of the key's balance is not held +by a sibling reservation. Only provider types we explicitly approve may do +that; everything else settles on token pricing. +""" + +import ast +import logging +import os +from pathlib import Path +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 import cost_calculation +from routstr.payment.cost_calculation import CostData, calculate_cost + +# 1000 input + 500 output tokens at the fixture rates below. +TOKEN_PRICED_MSATS = 20_000 +# 1.0 USD at the patched sats price, before any provider fee. +REPORTED_COST_MSATS = 20_000_000 + + +@pytest.fixture(autouse=True) +def fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 10) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 20) + + +@pytest.fixture(autouse=True) +def patch_sats_usd_price() -> None: # type: ignore[misc] + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5): + yield + + +@pytest.fixture(autouse=True) +def unit_provider_fee() -> None: # type: ignore[misc] + with patch( + "routstr.payment.cost_calculation._resolve_provider_fee", return_value=1.0 + ): + yield + + +def _response(**usage_extra: object) -> dict: + return { + "model": "gpt-4", + "usage": {"prompt_tokens": 1000, "completion_tokens": 500, **usage_extra}, + } + + +@pytest.fixture +def cost_log(caplog: pytest.LogCaptureFixture) -> pytest.LogCaptureFixture: # type: ignore[misc] + logger = logging.getLogger("routstr.payment.cost_calculation") + logger.addHandler(caplog.handler) + caplog.set_level(logging.WARNING) + yield caplog + logger.removeHandler(caplog.handler) + + +def _warned(caplog: pytest.LogCaptureFixture, needle: str) -> bool: + return any( + needle in rec.getMessage() + for rec in caplog.records + if rec.levelno >= logging.WARNING + ) + + +@pytest.mark.asyncio +async def test_untrusted_provider_falls_back_to_token_pricing() -> None: + result = await calculate_cost(_response(cost=1.0), max_cost=100_000) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + + +@pytest.mark.asyncio +async def test_trusted_provider_bills_the_reported_cost() -> None: + result = await calculate_cost( + _response(cost=1.0), max_cost=100_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert result.total_msats == REPORTED_COST_MSATS + + +@pytest.mark.asyncio +async def test_untrusted_huge_reported_cost_cannot_exceed_token_pricing() -> None: + result = await calculate_cost(_response(cost=1e9), max_cost=100_000) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + + +@pytest.mark.asyncio +async def test_untrusted_reported_cost_is_rejected_at_the_root() -> None: + response = _response() + response["cost"] = 1.0 + + result = await calculate_cost(response, max_cost=100_000) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + + +@pytest.mark.asyncio +async def test_untrusted_reported_cost_is_rejected_in_cost_details() -> None: + result = await calculate_cost( + _response(cost_details={"total_cost": 1.0}), max_cost=100_000 + ) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + + +@pytest.mark.asyncio +async def test_trusted_reported_cost_is_honoured_in_cost_details() -> None: + result = await calculate_cost( + _response(cost_details={"total_cost": 1.0}), + max_cost=100_000, + trusts_reported_cost=True, + ) + + assert isinstance(result, CostData) + assert result.total_msats == REPORTED_COST_MSATS + + +@pytest.mark.parametrize( + "reported", + [-1.0, float("nan"), float("inf"), float("-inf"), "nan", "inf", "-inf"], +) +def test_coerce_usd_rejects_non_finite_and_negative(reported: object) -> None: + """Infinity must be rejected at coercion, not survive into ``math.ceil`` + and get swallowed by the USD path's broad exception handler.""" + assert cost_calculation._coerce_usd(reported) == 0.0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "reported", + [-1.0, float("nan"), float("inf"), float("-inf"), "nan", "inf", "-inf"], +) +async def test_non_finite_or_negative_reported_cost_is_ignored( + reported: object, +) -> None: + result = await calculate_cost( + _response(cost=reported), max_cost=100_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + + +@pytest.mark.asyncio +async def test_over_reported_cost_is_flagged( + cost_log: pytest.LogCaptureFixture, +) -> None: + result = await calculate_cost( + _response(cost=1.0), max_cost=100_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert result.total_msats == REPORTED_COST_MSATS + assert _warned(cost_log, "implausible against token pricing") + + +@pytest.mark.asyncio +async def test_under_reported_cost_is_flagged( + cost_log: pytest.LogCaptureFixture, +) -> None: + result = await calculate_cost( + _response(cost=0.00000001), max_cost=100_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert result.total_msats < TOKEN_PRICED_MSATS + assert _warned(cost_log, "implausible against token pricing") + + +@pytest.mark.asyncio +async def test_cost_far_above_the_reservation_is_flagged( + cost_log: pytest.LogCaptureFixture, +) -> None: + result = await calculate_cost( + _response(cost=1.0), max_cost=1_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert _warned(cost_log, "far above the reservation") + + +@pytest.mark.asyncio +async def test_plausible_reported_cost_is_not_flagged( + cost_log: pytest.LogCaptureFixture, +) -> None: + # 20_000 msats — exactly what token pricing would charge. + result = await calculate_cost( + _response(cost=0.001), max_cost=100_000, trusts_reported_cost=True + ) + + assert isinstance(result, CostData) + assert result.total_msats == TOKEN_PRICED_MSATS + assert not _warned(cost_log, "implausible against token pricing") + assert not _warned(cost_log, "far above the reservation") + + +@pytest.mark.asyncio +async def test_ppq_byok_still_bills_above_the_reservation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The PPQ.AI BYOK payload from issue #615 legitimately settles ~9x the + reservation; the trust policy must not clamp it.""" + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001) + response = { + "model": "glm-5.2-fast", + "usage": { + "prompt_tokens": 164371, + "completion_tokens": 99, + "cost": 0.002260057305, + "is_byok": True, + "prompt_tokens_details": {"cached_tokens": 159301}, + "cost_details": { + "upstream_inference_cost": 0.04475361, + "upstream_inference_prompt_cost": 0.04410021, + "upstream_inference_completions_cost": 0.0006534, + }, + }, + } + + result = await calculate_cost(response, max_cost=100_000, trusts_reported_cost=True) + + assert isinstance(result, CostData) + assert result.total_msats == 940274 + + +def test_base_provider_does_not_trust_reported_cost() -> None: + from routstr.upstream.base import BaseUpstreamProvider + from routstr.upstream.generic import GenericUpstreamProvider + from routstr.upstream.routstr import RoutstrUpstreamProvider + + assert BaseUpstreamProvider.trusts_reported_cost is False + assert GenericUpstreamProvider.trusts_reported_cost is False + assert RoutstrUpstreamProvider.trusts_reported_cost is False + + +def test_approved_provider_types_trust_reported_cost() -> None: + from routstr.upstream.openrouter import OpenRouterUpstreamProvider + from routstr.upstream.ppqai import PPQAIUpstreamProvider + + assert OpenRouterUpstreamProvider.trusts_reported_cost is True + assert PPQAIUpstreamProvider.trusts_reported_cost is True + + +def test_every_settlement_site_passes_the_trust_flag() -> None: + """Streaming and non-streaming settlement must agree: a path that forgets + the flag silently reverts to trusting whatever the upstream reported.""" + source = Path(cost_calculation.__file__).parent.parent / "upstream" / "base.py" + tree = ast.parse(source.read_text()) + + missing = [ + node.lineno + for node in ast.walk(tree) + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and node.func.id in ("adjust_payment_for_tokens", "calculate_cost") + and not any(kw.arg == "trusts_reported_cost" for kw in node.keywords) + ] + + assert missing == [] diff --git a/tests/unit/test_settlement_identity.py b/tests/unit/test_settlement_identity.py index 6febb32d..a749e8da 100644 --- a/tests/unit/test_settlement_identity.py +++ b/tests/unit/test_settlement_identity.py @@ -20,9 +20,7 @@ from routstr.payment.cost_calculation import CostData, calculate_cost from routstr.payment.models import Architecture, Model, Pricing -def _make_model( - model_id: str, prompt_sats: float, completion_sats: float -) -> Model: +def _make_model(model_id: str, prompt_sats: float, completion_sats: float) -> Model: return Model( id=model_id, name=model_id, @@ -56,18 +54,14 @@ RESPONSE = { @pytest.fixture(autouse=True) def patch_sats_usd_price() -> None: # type: ignore[misc] - with patch( - "routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4 - ): + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4): yield @pytest.mark.asyncio async def test_served_model_pricing_wins_over_alias_lookup() -> None: """With ``model_obj`` given, the alias map is not consulted for pricing.""" - with patch( - "routstr.proxy.get_model_instance", return_value=WINNER - ) as alias_lookup: + with patch("routstr.proxy.get_model_instance", return_value=WINNER) as alias_lookup: result = await calculate_cost( dict(RESPONSE), max_cost=100_000, model_obj=SERVED ) @@ -97,11 +91,13 @@ async def test_usd_cost_path_applies_given_provider_fee() -> None: response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type] best_ranked = Mock(provider_fee=1.0) - with patch( - "routstr.proxy.get_provider_for_model", return_value=[best_ranked] - ): + with patch("routstr.proxy.get_provider_for_model", return_value=[best_ranked]): result = await calculate_cost( - response, max_cost=100_000, model_obj=SERVED, provider_fee=1.5 + response, + max_cost=100_000, + model_obj=SERVED, + provider_fee=1.5, + trusts_reported_cost=True, ) assert isinstance(result, CostData) @@ -118,10 +114,10 @@ async def test_usd_cost_path_falls_back_to_best_ranked_fee() -> None: response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type] best_ranked = Mock(provider_fee=2.0) - with patch( - "routstr.proxy.get_provider_for_model", return_value=[best_ranked] - ): - result = await calculate_cost(response, max_cost=100_000) + with patch("routstr.proxy.get_provider_for_model", return_value=[best_ranked]): + result = await calculate_cost( + response, max_cost=100_000, trusts_reported_cost=True + ) assert isinstance(result, CostData) assert result.total_msats == 4_000 @@ -140,9 +136,7 @@ async def test_x_cashu_cost_uses_served_model_not_upstream_echo() -> None: provider = GenericUpstreamProvider("http://upstream.example", "key", 1.0) response = dict(RESPONSE, model="totally-unknown-wire-name") - with patch( - "routstr.proxy.get_model_instance", return_value=WINNER - ) as alias_lookup: + with patch("routstr.proxy.get_model_instance", return_value=WINNER) as alias_lookup: cost = await provider.get_x_cashu_cost( response, max_cost_for_model=100_000, model_obj=SERVED )