fix: gate upstream-reported cost behind a per-provider trust policy

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.
This commit is contained in:
9qeklajc
2026-08-25 00:38:12 +02:00
parent 7f4be83b14
commit a7d4e2832d
10 changed files with 528 additions and 112 deletions
+11 -1
View File
@@ -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):
+127 -24
View File
@@ -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",
+20
View File
@@ -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(
+4
View File
@@ -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.
+4
View File
@@ -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__(
+16 -2
View File
@@ -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
+25 -27
View File
@@ -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.
+30 -38
View File
@@ -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)
+277
View File
@@ -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 == []
+14 -20
View File
@@ -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
)