mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: compare effective cache rates when flooring pinned endpoints at an override
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user