From 057a752b1e908b8b88a734cc6edd7d9d820e954b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 12 Jul 2026 13:07:05 +0200 Subject: [PATCH] fix cost breakdown calculation --- routstr/payment/cost_calculation.py | 127 ++++++++++++++------ tests/unit/test_cost_calculation_caching.py | 104 ++++++++++++++++ 2 files changed, 191 insertions(+), 40 deletions(-) diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 84f3847f..25e94993 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -157,11 +157,16 @@ async def calculate_cost( }, ) try: + cost_details = usage_data.get("cost_details", {}) + if not isinstance(cost_details, dict): + cost_details = {} input_usd = _coerce_usd( - usage_data.get("cost_details", {}).get("input_cost", 0) + cost_details.get("input_cost") + or cost_details.get("upstream_inference_prompt_cost") ) output_usd = _coerce_usd( - usage_data.get("cost_details", {}).get("output_cost", 0) + cost_details.get("output_cost") + or cost_details.get("upstream_inference_completions_cost") ) return _calculate_from_usd_cost( usd_cost, @@ -281,54 +286,92 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: def _get_pricing_rates( response_data: dict, ) -> tuple[float, float, float, float] | None: - """Get model-based pricing rates or None if using fixed pricing. + """Get configured rates, falling back to LiteLLM's model cost map. - Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate) + Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate). + ``None`` means configured fixed pricing should be used by the caller. """ - if settings.fixed_pricing: + if settings.fixed_pricing and ( + settings.fixed_per_1k_input_tokens + or settings.fixed_per_1k_output_tokens + ): return None from ..proxy import get_model_instance + from .models import litellm_cost_entry response_model = response_data.get("model", "") model_obj = get_model_instance(response_model) - if not model_obj: - logger.error("Invalid model in response", extra={"response_model": response_model}) - raise ValueError(f"Invalid model: {response_model}") + if model_obj and model_obj.sats_pricing: + try: + mspp = float(model_obj.sats_pricing.prompt) + mspc = float(model_obj.sats_pricing.completion) + mscr = float(model_obj.sats_pricing.input_cache_read or 0) + mscw = float(model_obj.sats_pricing.input_cache_write or 0) - if not model_obj.sats_pricing: - logger.error( - "Model pricing not defined", - extra={"model": response_model, "model_id": response_model}, + mspp_1k = mspp * 1_000_000.0 + mspc_1k = mspc * 1_000_000.0 + mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k + mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k + source = "configured" + except Exception as e: + logger.error("Invalid pricing data", extra={"error": str(e)}) + raise ValueError("Invalid pricing data") from e + else: + pricing_model = ( + model_obj.forwarded_model_id if model_obj else None + ) or response_model + pricing = litellm_cost_entry(pricing_model) + if pricing is None: + logger.error( + "Model pricing not found in configured models or LiteLLM", + extra={ + "response_model": response_model, + "pricing_model": pricing_model, + }, + ) + raise ValueError(f"Pricing not found for model: {response_model}") + + input_usd = _coerce_usd(pricing.get("input_cost_per_token")) + output_usd = _coerce_usd(pricing.get("output_cost_per_token")) + if input_usd <= 0 or output_usd <= 0: + raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}") + + provider_fee = _resolve_provider_fee(response_model) + usd_per_sat = sats_usd_price() + mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat + mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat + cache_read_usd = _coerce_usd( + pricing.get("cache_read_input_token_cost") ) - raise ValueError("Model pricing not defined") - - try: - mspp = float(model_obj.sats_pricing.prompt) - mspc = float(model_obj.sats_pricing.completion) - mscr = float(model_obj.sats_pricing.input_cache_read or 0) - mscw = float(model_obj.sats_pricing.input_cache_write or 0) - - mspp_1k = mspp * 1_000_000.0 - mspc_1k = mspc * 1_000_000.0 - mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k - mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k - - logger.info( - "Applied model-specific pricing", - extra={ - "model": response_model, - "input_price_msats_per_1k": mspp_1k, - "output_price_msats_per_1k": mspc_1k, - "cache_read_price_msats_per_1k": mscr_1k, - "cache_write_price_msats_per_1k": mscw_1k, - }, + cache_write_usd = _coerce_usd( + pricing.get("cache_creation_input_token_cost") ) - return mspp_1k, mspc_1k, mscr_1k, mscw_1k - except Exception as e: - logger.error("Invalid pricing data", extra={"error": str(e)}) - raise ValueError("Invalid pricing data") from e + mscr_1k = ( + cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat + if cache_read_usd > 0 + else mspp_1k + ) + mscw_1k = ( + cache_write_usd * provider_fee * 1_000_000.0 / usd_per_sat + if cache_write_usd > 0 + else mspp_1k + ) + source = "litellm" + + logger.info( + "Applied model-specific pricing", + extra={ + "model": response_model, + "pricing_source": source, + "input_price_msats_per_1k": mspp_1k, + "output_price_msats_per_1k": mspc_1k, + "cache_read_price_msats_per_1k": mscr_1k, + "cache_write_price_msats_per_1k": mscw_1k, + }, + ) + return mspp_1k, mspc_1k, mscr_1k, mscw_1k def _resolve_provider_fee(model_id: str) -> float: @@ -367,8 +410,12 @@ def _calculate_from_usd_cost( cost_in_msats = math.ceil(cost_in_sats * 1000) if input_usd > 0 or output_usd > 0: - input_msats = int((input_usd * sats_per_usd) * 1000) - output_msats = int((output_usd * sats_per_usd) * 1000) + # The total is the authoritative billed amount. Allocating that integer + # total proportionally avoids losing sub-millisatoshi remainders when + # input and output components are each truncated independently. + component_usd = input_usd + output_usd + input_msats = math.floor(cost_in_msats * input_usd / component_usd) + output_msats = cost_in_msats - input_msats else: effective_input_tokens = ( input_tokens + cache_read_tokens + cache_creation_tokens diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 93c68877..6ff7170b 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -465,9 +465,113 @@ async def test_cache_read_only_usd_cost_response_is_billed( assert result.cache_read_input_tokens == 1000 +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("total_cost", "input_cost", "output_cost", "expected_msats"), + [ + (0.000471, 0.00023451, 0.00023649, 9420), + (0.00000004, 0.00000002, 0.00000002, 1), + ], +) +async def test_small_usd_cost_components_sum_to_rounded_total( + total_cost: float, + input_cost: float, + output_cost: float, + expected_msats: int, +) -> None: + """Small USD component costs must retain every billed millisatoshi.""" + response = { + "model": "gpt-4", + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "cost_details": { + "total_cost": total_cost, + "input_cost": input_cost, + "output_cost": output_cost, + }, + }, + } + + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + assert result.total_msats == expected_msats + assert result.input_msats + result.output_msats == result.total_msats + + +@pytest.mark.asyncio +async def test_openrouter_upstream_inference_cost_components_are_used() -> None: + """OpenRouter component aliases must determine the input/output split.""" + response = { + "model": "gpt-4", + "usage": { + "prompt_tokens": 375, + "completion_tokens": 158, + "total_tokens": 533, + "cost": 0.00022354, + "is_byok": False, + "prompt_tokens_details": { + "cached_tokens": 286, + "cache_write_tokens": 0, + }, + "cost_details": { + "upstream_inference_cost": 0.00022354, + "upstream_inference_prompt_cost": 0.00004974, + "upstream_inference_completions_cost": 0.0001738, + }, + "completion_tokens_details": {"reasoning_tokens": 17}, + }, + } + + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + assert result.input_msats == 994 + assert result.output_msats == 3477 + assert result.input_msats + result.output_msats == result.total_msats == 4471 + + # ============================================================================ # Test 13: Missing Usage Block # ============================================================================ +@pytest.mark.asyncio +async def test_missing_upstream_cost_uses_litellm_model_pricing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Token usage without an upstream cost is priced from LiteLLM.""" + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0) + monkeypatch.setattr( + "routstr.payment.models.litellm_cost_entry", + lambda model: { + "input_cost_per_token": 0.000001, + "output_cost_per_token": 0.000002, + }, + ) + response = { + "model": "priced-by-litellm", + "usage": { + "prompt_tokens": 90, + "completion_tokens": 80, + "total_tokens": 170, + "prompt_tokens_details": {"cached_tokens": 0}, + "completion_tokens_details": {"reasoning_tokens": 74}, + "prompt_cache_hit_tokens": 0, + "prompt_cache_miss_tokens": 90, + }, + } + + result = await calculate_cost(response, max_cost=10000) + + assert isinstance(result, CostData) + assert not isinstance(result, MaxCostData) + assert result.input_msats == 1800 + assert result.output_msats == 3200 + assert result.input_msats + result.output_msats == result.total_msats == 5000 + + @pytest.mark.asyncio async def test_missing_usage_block(mock_fixed_pricing: None) -> None: """When usage is missing, return MaxCostData with zero tokens."""