diff --git a/routstr/core/admin.py b/routstr/core/admin.py index f6f24db6..a36aa7bf 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1727,16 +1727,14 @@ async def certify_upstream_provider( # The proxy reserves and token-bills a pinned endpoint at that # endpoint's own rates, so the cost rows price it the same way. if selected_path is not None: - from ..upstream.model_paths import price_pinned_endpoint + from ..upstream.model_paths import apply_model_path_pricing - async with create_session() as session: - model_obj = await price_pinned_endpoint( - session, - model_obj, - selected_path, - provider.provider_fee, - sats_to_usd, - ) + model_obj = apply_model_path_pricing( + model_obj, + selected_path, + provider.provider_fee, + sats_to_usd, + ) # The timeout applies per upstream call. The run makes up to six # calls (models, two short probes after a max_completion_tokens retry, # three cache probes), so the request can stay open for six times it. diff --git a/routstr/proxy.py b/routstr/proxy.py index 7c210e0c..9faa1acb 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -53,9 +53,9 @@ from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request from .upstream.helpers import init_upstreams from .upstream.model_paths import ( ModelPathSelector, + apply_model_path_pricing, decode_model_path, is_openrouter_base_url, - price_pinned_endpoint, public_model_id, public_provider_url, ) @@ -194,9 +194,7 @@ async def _price_pinned_endpoint( extra={"model": selector.model_id, "endpoint": selector.endpoint_tag}, ) return model_obj - return await price_pinned_endpoint( - session, model_obj, row, upstream.provider_fee, sats_to_usd - ) + return apply_model_path_pricing(model_obj, row, upstream.provider_fee, sats_to_usd) def get_model_instance(model_id: str) -> Model | None: diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index a72fe807..b69bef9e 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -38,7 +38,7 @@ from ..core.logging import get_logger if TYPE_CHECKING: from sqlmodel.ext.asyncio.session import AsyncSession - from ..payment.models import Model, Pricing + from ..payment.models import Model from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -893,7 +893,6 @@ def apply_model_path_pricing( row: ModelPathRow, provider_fee: float, sats_to_usd: float, - floor: "Pricing | None" = None, ) -> "Model": """Return ``model`` priced from an exact endpoint path's own rates. @@ -903,9 +902,6 @@ def apply_model_path_pricing( default listing; OpenRouter charges the endpoint that serves the request, so the proxy reserves and token-bills a pinned endpoint with them, using the same limits ``/v1/models/paths`` quotes its max cost from. - - ``floor`` is an operator price override: no rate is billed below it, so - pinning an endpoint cannot bypass the operator's pricing. """ if row.endpoint_tag is None: return model @@ -928,20 +924,9 @@ def apply_model_path_pricing( model.forwarded_model_id or row.model_id, Pricing.parse_obj(metadata["pricing"]), ) - rates = { - key: float(value) * provider_fee - for key, value in pricing.dict().items() - 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, floor_rates[key]) - for key, value in _effective_cache_rates(rates).items() - } - pricing = Pricing.parse_obj(rates) + pricing = Pricing.parse_obj( + {key: float(value) * provider_fee for key, value in pricing.dict().items()} + ) update: dict[str, Any] = {"pricing": pricing, "sats_pricing": None} context_length = metadata.get("context_length") max_completion_tokens = metadata.get("max_completion_tokens") @@ -966,51 +951,6 @@ 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", - row: ModelPathRow, - provider_fee: float, - sats_to_usd: float, -) -> "Model": - """Price ``model`` for a request pinned to ``row``'s endpoint. - - An enabled operator override for the model on this provider is the - price floor; ``model`` already carries it, since overrides replace the - provider's model in routing. - """ - override = ( - await session.exec( - select(ModelRow).where( - ModelRow.id == model.id, - ModelRow.upstream_provider_id == row.upstream_provider_id, - ModelRow.enabled, - ) - ) - ).first() - return apply_model_path_pricing( - model, - row, - provider_fee, - sats_to_usd, - floor=model.pricing if override is not None else None, - ) - - def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]: endpoint = None if row.endpoint_tag or row.endpoint_name: diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 8d338850..065f888f 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -409,16 +409,13 @@ async def test_certify_margin_bills_pinned_path_pricing( assert resp.status_code == 200, resp.text margin = _find_row(resp.json()["rows"], "cost.margin") - # The proxy bills a pinned endpoint at its own rates, with the operator's - # override as the floor per rate: the endpoint's prompt rate, the - # override's completion and cache-read rates. The model's rates alone - # would give 3, 356 and 43; the endpoint's alone 3, 289 and 15, which - # misses the cached call's reported cost of 26. + # The proxy token-bills a pinned endpoint at its own rates, not the + # model's (3, 356, 43). Those rates miss the cached call's reported cost. assert [ (sample["upstream_msats_with_fee"], sample["configured_msats"]) for sample in margin["evidence"]["samples"] - ] == [(3, 3), (269, 289), (26, 27)] - assert margin["status"] == "ok" + ] == [(3, 3), (269, 289), (26, 15)] + assert margin["status"] == "fail" @pytest.mark.integration diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 0f12eef9..a62b0eeb 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -910,14 +910,9 @@ def _endpoint_row( ) -def _session_with_rows(rows: list[Any], override: Any = None) -> MagicMock: - """Answers the path-row query with ``rows``, the override one with ``override``.""" +def _session_with_rows(rows: list[Any]) -> MagicMock: session = MagicMock() - session.exec = AsyncMock( - return_value=MagicMock( - all=MagicMock(return_value=rows), first=MagicMock(return_value=override) - ) - ) + session.exec = AsyncMock(return_value=MagicMock(all=MagicMock(return_value=rows))) return session @@ -1018,71 +1013,3 @@ async def test_endpoint_pin_reserves_the_max_cost_paths_quotes() -> None: assert priced.sats_pricing is not None and priced.top_provider is not None assert priced.sats_pricing.max_cost == pytest.approx(quoted) assert priced.top_provider.max_completion_tokens == 8192 - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("override", "expected"), - [ - # The operator's override is above the endpoint: it stays the price. - ((6e-6, 12e-6), (6e-6, 12e-6)), - # The endpoint costs more than the override: never bill below cost. - ((1e-6, 2e-6), (2e-6, 4e-6)), - # Each rate takes the higher of the two. - ((3e-6, 1e-6), (3e-6, 4e-6)), - ], -) -async def test_endpoint_pin_never_bills_below_an_operator_override( - override: tuple[float, float], expected: tuple[float, float] -) -> None: - from routstr.core.db import ModelRow - - model = _priced_model(*override) - with patch("routstr.payment.price.SATS_USD_PRICE", _SATS_USD): - priced = await proxy_module._price_pinned_endpoint( - _session_with_rows([_endpoint_row()], override=MagicMock(spec=ModelRow)), - _endpoint_selector(), - model, - _openrouter_upstream(), - ) - - 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)