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 not key.startswith("max_")
|
||||||
}
|
}
|
||||||
if floor is not None:
|
if floor is not None:
|
||||||
|
floor_rates = _effective_cache_rates(
|
||||||
|
{key: float(getattr(floor, key)) for key in rates}
|
||||||
|
)
|
||||||
rates = {
|
rates = {
|
||||||
key: max(value, float(getattr(floor, key)))
|
key: max(value, floor_rates[key])
|
||||||
for key, value in rates.items()
|
for key, value in _effective_cache_rates(rates).items()
|
||||||
}
|
}
|
||||||
pricing = Pricing.parse_obj(rates)
|
pricing = Pricing.parse_obj(rates)
|
||||||
update: dict[str, Any] = {"pricing": pricing, "sats_pricing": None}
|
update: dict[str, Any] = {"pricing": pricing, "sats_pricing": None}
|
||||||
@@ -963,6 +966,20 @@ def apply_model_path_pricing(
|
|||||||
return model
|
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(
|
async def price_pinned_endpoint(
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
model: "Model",
|
model: "Model",
|
||||||
|
|||||||
@@ -853,7 +853,9 @@ _SATS_USD = 0.001
|
|||||||
_ENDPOINT_PRICING = {"prompt": 2e-6, "completion": 4e-6}
|
_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 (
|
from routstr.payment.models import (
|
||||||
Architecture,
|
Architecture,
|
||||||
Model,
|
Model,
|
||||||
@@ -875,7 +877,9 @@ def _priced_model(prompt: float = 1e-6, completion: float = 2e-6) -> Any:
|
|||||||
tokenizer="unknown",
|
tokenizer="unknown",
|
||||||
instruct_type=None,
|
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,
|
model.pricing.max_prompt_cost,
|
||||||
@@ -888,6 +892,7 @@ def _priced_model(prompt: float = 1e-6, completion: float = 2e-6) -> Any:
|
|||||||
def _endpoint_row(
|
def _endpoint_row(
|
||||||
model_id: str = MODEL_ID,
|
model_id: str = MODEL_ID,
|
||||||
endpoint_tag: str = "deepinfra/fp8",
|
endpoint_tag: str = "deepinfra/fp8",
|
||||||
|
pricing: dict[str, float] | None = None,
|
||||||
**limits: int,
|
**limits: int,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
from routstr.core.db import ModelPathRow
|
from routstr.core.db import ModelPathRow
|
||||||
@@ -899,7 +904,7 @@ def _endpoint_row(
|
|||||||
provider_type="openrouter",
|
provider_type="openrouter",
|
||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
model_metadata=json.dumps(
|
model_metadata=json.dumps(
|
||||||
{"id": model_id, "pricing": _ENDPOINT_PRICING, **limits}
|
{"id": model_id, "pricing": pricing or _ENDPOINT_PRICING, **limits}
|
||||||
),
|
),
|
||||||
upstream_provider_id=1,
|
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)
|
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