diff --git a/routstr/algorithm.py b/routstr/algorithm.py index cced61f7..68406876 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -135,9 +135,21 @@ def create_model_mappings( Returns: Tuple of (model_instances, provider_map, unique_models) """ - from .payment.models import _row_to_model + from .payment.models import _row_to_model, has_usable_pricing from .upstream.helpers import resolve_model_alias + def _unusable_price(model: "Model") -> bool: + """A candidate may only route on rates a request can be billed against. + + Mirrors the served-catalog backstop in ``list_models``: a negative or + non-finite rate is not a price, and the cost calculation cannot bill on + one, so every request on the model would be charged the full maximum + reservation instead. Applies to provider-discovered models as well as + persisted overrides — no override row need exist for a malformed price + to be built into the candidate map. + """ + return not has_usable_pricing(model.pricing) + candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} unique_models: dict[str, "Model"] = {} unique_model_keys: dict[str, str] = {} @@ -231,6 +243,9 @@ def create_model_mappings( else: model_to_use = model + if _unusable_price(model_to_use): + continue + forwarded_model_id = get_effective_forwarded_model_id(model_to_use) # Get all aliases for this model @@ -297,6 +312,8 @@ def create_model_mappings( continue if not model_to_use.enabled: continue + if _unusable_price(model_to_use): + continue forwarded_model_id = get_effective_forwarded_model_id(model_to_use) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 6595d684..23c7ab6f 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -97,6 +97,19 @@ def is_usable_rate(rate: float) -> bool: return math.isfinite(rate) and rate >= 0.0 +def has_usable_pricing(pricing: Pricing) -> bool: + """True if every billable rate is a number a request could be billed on. + + Free is usable — a rate of zero is a real price. This asks only whether the + price is well-formed. One unusable rate disqualifies the whole price even + alongside a valid one, since a request can bill on the bad field: a positive + ``completion`` must not hide a negative ``prompt``. + """ + return all( + is_usable_rate(getattr(pricing, field)) for field in BILLABLE_PRICING_FIELDS + ) + + class TopProvider(BaseModel): context_length: int | None = None max_completion_tokens: int | None = None @@ -354,21 +367,39 @@ async def list_models( rows = (await session.exec(query)).all() # type: ignore provider_result = await session.exec(select(UpstreamProviderRow)) providers_by_id = {p.id: p for p in provider_result.all()} - return [ - _row_to_model( + + models: list[Model] = [] + for r in rows: + if not include_disabled and not ( + r.upstream_provider_id in providers_by_id + and providers_by_id[r.upstream_provider_id].enabled + ): + continue + model = _row_to_model( r, apply_provider_fee=apply_fees, provider_fee=providers_by_id[r.upstream_provider_id].provider_fee if r.upstream_provider_id in providers_by_id else 1.01, ) - for r in rows - if include_disabled - or ( - r.upstream_provider_id in providers_by_id - and providers_by_id[r.upstream_provider_id].enabled - ) - ] + # Served-map backstop for legacy rows and writers that bypass the admin + # edge: a negative or non-finite rate is not a price. Serving one + # advertises a rate the cost calculation cannot bill on, so the request + # falls through to the flat maximum reservation — or, if the rate is + # negative, bills an amount settlement credits back to the caller. + # ``include_disabled`` is the operator's listing, which must keep showing + # the row so it can be repaired. + if not include_disabled and not has_usable_pricing(model.pricing): + logger.warning( + "Withholding model with an unusable stored rate from the catalog", + extra={ + "model_id": r.id, + "upstream_provider_id": r.upstream_provider_id, + }, + ) + continue + models.append(model) + return models def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]: diff --git a/tests/integration/test_served_catalog_rate_backstop.py b/tests/integration/test_served_catalog_rate_backstop.py new file mode 100644 index 00000000..96455554 --- /dev/null +++ b/tests/integration/test_served_catalog_rate_backstop.py @@ -0,0 +1,157 @@ +"""The served catalog is the last guard between a stored row and a charge. + +Stored pricing is JSON written by whatever produced the row — an upstream +import, an operator, a legacy migration, or a foreign writer that never passed +the admin edge. So the read path cannot assume a stored rate is a number: it +must decline to serve a row it cannot bill on, and it must survive a row it +cannot read at all rather than taking the whole catalog down with it. + +The admin listing is deliberately exempt: it includes disabled models and is the +one view that still shows the operator the row that needs repair. +""" + +from __future__ import annotations + +import json + +import pytest +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.payment.models import list_models +from routstr.proxy import reinitialize_upstreams + +_ARCHITECTURE = json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } +) + + +async def _make_provider(session: AsyncSession) -> int: + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://served-upstream.example/v1", + api_key="test-key", + provider_fee=1.0, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + await reinitialize_upstreams() + assert provider.id is not None + return provider.id + + +async def _insert_row( + session: AsyncSession, + provider_id: int, + *, + model_id: str, + pricing: dict[str, object], +) -> None: + session.add( + ModelRow( + id=model_id, + name=model_id, + description="d", + created=0, + context_length=8192, + architecture=_ARCHITECTURE, + pricing=json.dumps(pricing), + upstream_provider_id=provider_id, + enabled=True, + forwarded_model_id=model_id, + ) + ) + await session.commit() + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_rate", + [float("nan"), float("inf"), -1.0], + ids=["nan", "inf", "negative"], +) +async def test_served_catalog_excludes_a_malformed_stored_rate( + integration_session: AsyncSession, bad_rate: float +) -> None: + """A stored rate that is not a number must not be advertised. + + Zero is a real price and a free model is servable, but a negative or + non-finite rate is not a price at all: serving it advertises a rate the cost + calculation cannot bill on, so every request falls through to the flat + maximum reservation — or, for a negative rate, bills an amount settlement + credits back to the caller. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="good", + pricing={"prompt": 1e-06, "completion": 2e-06}, + ) + await _insert_row( + integration_session, + provider_id, + model_id="bad-rate", + pricing={"prompt": bad_rate, "completion": 2e-06}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"good"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_free_stored_price_is_still_served( + integration_session: AsyncSession, +) -> None: + """Zero is a real price. Rejecting malformed rates must not also drop a row + priced at zero, which is a free model and not a broken one.""" + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="free", + pricing={"prompt": 0.0, "completion": 0.0}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"free"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_listing_still_shows_a_malformed_stored_rate( + integration_session: AsyncSession, +) -> None: + """The operator has to be able to see the row that needs fixing. + + The backstop keeps a malformed row out of the *served* catalog. The listing + that includes disabled models is the one view where the row must still + appear, or the operator loses the ability to repair it. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="bad-rate", + pricing={"prompt": -1.0, "completion": 2e-06}, + ) + + listed = { + m.id + for m in await list_models( + integration_session, provider_id, include_disabled=True + ) + } + + assert listed == {"bad-rate"} diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index 5d31af2e..62d3ad0d 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -954,3 +954,67 @@ def test_create_model_mappings_uppercase_prefixed_base_keeps_top_tier() -> None: assert provider_map["qwen2.5-72b"][0] == (prefixed_cheap, prefixed_provider) assert unique_models["qwen2.5-72b"].upstream_provider_id == "prefixed" + + +def test_create_model_mappings_excludes_a_malformed_price() -> None: + """A rate that is not a number must not be routable. + + A negative or non-finite rate reads as a real price to every truthiness + check, so the candidate was built into the map and served. The cost + calculation cannot price on such a rate, so every request on the model fell + through to the flat maximum reservation — or, for a negative rate, billed a + negative amount that settlement credits back to the caller. + """ + healthy = create_test_model("healthy-model") + for bad_rate in (float("nan"), float("inf"), -1.0): + broken = create_test_model("broken-model", prompt_price=bad_rate) + provider = create_test_provider( + "custom", + "https://custom.example/v1", + db_id=1, + models=[broken, healthy], + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_key={}, + disabled_model_keys=set(), + ) + + assert "broken-model" not in provider_map, bad_rate + assert "broken-model" not in unique_models, bad_rate + # One unroutable candidate must not cost the provider its other models. + assert "healthy-model" in provider_map, bad_rate + + +def test_create_model_mappings_excludes_an_override_with_a_malformed_price( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An override row carrying a malformed rate is unroutable too. + + An override replaces the discovered model's price, so a provider whose + catalog is sound still routes at whatever the row says. The guard has to sit + after the override is applied, not before it. + """ + discovered = create_test_model("shared-model") + provider = create_test_provider( + "custom", "https://custom.example/v1", db_id=3, models=[discovered] + ) + override_model = create_test_model("shared-model", prompt_price=float("-inf")) + + monkeypatch.setattr( + "routstr.payment.models._row_to_model", + lambda *args, **kwargs: override_model, + ) + override_row = SimpleNamespace( + id="shared-model", upstream_provider_id=3, enabled=True + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_key={("shared-model", 3): (override_row, 1.0)}, + disabled_model_keys=set(), + ) + + assert "shared-model" not in provider_map + assert "shared-model" not in unique_models