diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 581cf6e4..a72fe807 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -934,9 +934,12 @@ def apply_model_path_pricing( if not key.startswith("max_") } if floor is not None: + floor_rates = _effective_cache_rates( + {key: float(getattr(floor, key)) for key in rates} + ) rates = { - key: max(value, float(getattr(floor, key))) - for key, value in rates.items() + key: max(value, floor_rates[key]) + for key, value in _effective_cache_rates(rates).items() } pricing = Pricing.parse_obj(rates) update: dict[str, Any] = {"pricing": pricing, "sats_pricing": None} @@ -963,6 +966,20 @@ def apply_model_path_pricing( return model +def _effective_cache_rates(rates: dict[str, float]) -> dict[str, float]: + """Spell out cache rates settlement reads as "bill at the prompt rate". + + A zero cache rate is billed at the prompt rate, so a per-rate ``max`` + must compare those prompt rates, not the zeros. + """ + return { + key: value + if value > 0 or key not in ("input_cache_read", "input_cache_write") + else rates["prompt"] + for key, value in rates.items() + } + + async def price_pinned_endpoint( session: AsyncSession, model: "Model", diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 765ed216..0f12eef9 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -853,7 +853,9 @@ _SATS_USD = 0.001 _ENDPOINT_PRICING = {"prompt": 2e-6, "completion": 4e-6} -def _priced_model(prompt: float = 1e-6, completion: float = 2e-6) -> Any: +def _priced_model( + prompt: float = 1e-6, completion: float = 2e-6, cache_read: float = 0.0 +) -> Any: from routstr.payment.models import ( Architecture, Model, @@ -875,7 +877,9 @@ def _priced_model(prompt: float = 1e-6, completion: float = 2e-6) -> Any: tokenizer="unknown", instruct_type=None, ), - pricing=Pricing(prompt=prompt, completion=completion), + pricing=Pricing( + prompt=prompt, completion=completion, input_cache_read=cache_read + ), ) ( model.pricing.max_prompt_cost, @@ -888,6 +892,7 @@ def _priced_model(prompt: float = 1e-6, completion: float = 2e-6) -> Any: def _endpoint_row( model_id: str = MODEL_ID, endpoint_tag: str = "deepinfra/fp8", + pricing: dict[str, float] | None = None, **limits: int, ) -> Any: from routstr.core.db import ModelPathRow @@ -899,7 +904,7 @@ def _endpoint_row( provider_type="openrouter", endpoint_tag=endpoint_tag, model_metadata=json.dumps( - {"id": model_id, "pricing": _ENDPOINT_PRICING, **limits} + {"id": model_id, "pricing": pricing or _ENDPOINT_PRICING, **limits} ), upstream_provider_id=1, ) @@ -1042,3 +1047,42 @@ async def test_endpoint_pin_never_bills_below_an_operator_override( ) assert (priced.pricing.prompt, priced.pricing.completion) == pytest.approx(expected) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("override", "endpoint", "expected_cache_read"), + [ + # The override has no cache rate, so its cache reads bill at its + # prompt rate; the endpoint's cheaper cache rate must not undercut it. + ( + (2e-6, 8e-6, 0.0), + {"prompt": 1e-6, "completion": 2e-6, "input_cache_read": 1e-7}, + 2e-6, + ), + # The endpoint has no cache rate, so it charges its prompt rate; the + # override's cheaper cache rate must not bill below that cost. + ((1e-6, 2e-6, 1e-8), {"prompt": 2e-6, "completion": 4e-6}, 2e-6), + ], + ids=["override-without-cache-rate", "endpoint-without-cache-rate"], +) +async def test_override_floor_compares_effective_cache_rates( + override: tuple[float, float, float], + endpoint: dict[str, float], + expected_cache_read: float, +) -> None: + """A zero cache rate is billed at the prompt rate, so the floor compares + those effective rates rather than the zeros.""" + from routstr.core.db import ModelRow + + with patch("routstr.payment.price.SATS_USD_PRICE", _SATS_USD): + priced = await proxy_module._price_pinned_endpoint( + _session_with_rows( + [_endpoint_row(pricing=endpoint)], override=MagicMock(spec=ModelRow) + ), + _endpoint_selector(), + _priced_model(*override), + _openrouter_upstream(), + ) + + assert priced.pricing.input_cache_read == pytest.approx(expected_cache_read)