fix: compare effective cache rates when flooring pinned endpoints at an override

This commit is contained in:
9qeklajc
2026-10-03 17:28:46 +02:00
parent a09a5f1b38
commit 9068d681d7
2 changed files with 66 additions and 5 deletions
+19 -2
View File
@@ -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",
+47 -3
View File
@@ -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)