fix(models): don't trust foreign catalog pricing for custom/unknown providers

Custom providers (custom:xxx) serve their own pricing; models.dev stores
OpenRouter prices for the same model ids. The cost guard fired on that
foreign pricing and blocked composer/CLI model switches on custom
providers with a wildly wrong warning (#54348).

expensive_model_warning now only trusts model_info/models.dev pricing
when the provider maps to a models.dev provider and the info's
provider_id matches, and only consults the pricing-entry lookup when
the billing route is known. Salvaged from #54422; the PR's desktop-hook
half predates the use-model-controls rewrite and is superseded by the
hook's existing rollback handling.
This commit is contained in:
Dustin Persek
2026-08-14 01:14:08 -07:00
committed by Teknium
parent 9a8af40192
commit 54cc39aa15
2 changed files with 155 additions and 6 deletions
+53 -4
View File
@@ -6,7 +6,7 @@ from dataclasses import dataclass
from decimal import Decimal, InvalidOperation
from typing import Optional
from agent.models_dev import ModelInfo
from agent.models_dev import ModelInfo, PROVIDER_TO_MODELS_DEV
INPUT_COST_WARNING_THRESHOLD = Decimal("20")
@@ -54,6 +54,40 @@ def _pricing_from_model_info(
)
def _known_models_dev_provider(provider: Optional[str]) -> Optional[str]:
normalized = (provider or "").strip().lower()
if not normalized:
return None
return PROVIDER_TO_MODELS_DEV.get(normalized)
def _can_trust_model_info_pricing(
provider: Optional[str],
model_info: Optional[ModelInfo],
) -> bool:
expected_provider = _known_models_dev_provider(provider)
if not expected_provider or model_info is None:
return False
actual_provider = str(getattr(model_info, "provider_id", "") or "").strip().lower()
return not actual_provider or actual_provider == expected_provider
def _can_trust_pricing_lookup(
model_name: str,
*,
provider: Optional[str],
base_url: Optional[str],
) -> bool:
try:
from agent.usage_pricing import resolve_billing_route
route = resolve_billing_route(model_name, provider=provider, base_url=base_url)
except Exception:
return False
return route.billing_mode != "unknown"
def expensive_model_warning(
model_name: str,
*,
@@ -71,8 +105,18 @@ def expensive_model_warning(
if not model:
return None
input_cost, output_cost, source = _pricing_from_model_info(model_info)
if input_cost is None and output_cost is None and provider:
input_cost: Optional[Decimal] = None
output_cost: Optional[Decimal] = None
source = ""
if _can_trust_model_info_pricing(provider, model_info):
input_cost, output_cost, source = _pricing_from_model_info(model_info)
if (
input_cost is None
and output_cost is None
and _known_models_dev_provider(provider)
):
try:
from agent.models_dev import get_model_info
@@ -81,7 +125,12 @@ def expensive_model_warning(
)
except Exception:
pass
if input_cost is None and output_cost is None:
if (
input_cost is None
and output_cost is None
and _can_trust_pricing_lookup(model, provider=provider, base_url=base_url)
):
try:
from agent.usage_pricing import get_pricing_entry
+102 -2
View File
@@ -1,5 +1,7 @@
from decimal import Decimal
import pytest
from agent.models_dev import ModelInfo
from agent.usage_pricing import PricingEntry
from hermes_cli.model_cost_guard import expensive_model_warning
@@ -10,16 +12,114 @@ def test_no_warning_when_known_prices_are_at_threshold():
id="edge/model",
name="edge/model",
family="",
provider_id="test",
provider_id="anthropic",
cost_input=20.0,
cost_output=100.0,
)
assert expensive_model_warning("edge/model", provider="test", model_info=info) is None
assert expensive_model_warning("edge/model", provider="anthropic", model_info=info) is None
def test_warns_when_models_dev_input_price_exceeds_threshold():
info = ModelInfo(
id="expensive/input",
name="expensive/input",
family="",
provider_id="anthropic",
cost_input=20.01,
cost_output=1.0,
)
warning = expensive_model_warning(
"expensive/input",
provider="anthropic",
model_info=info,
)
assert warning is not None
assert warning.input_cost_per_million == Decimal("20.01")
assert "EXPENSIVE MODEL WARNING" in warning.message
assert "$20/M input" in warning.message
@pytest.mark.parametrize("provider", ["custom", "custom:routerai", "routerai"])
def test_skips_foreign_models_dev_pricing_for_custom_or_unknown_providers(provider):
info = ModelInfo(
id="openai/gpt-5.5-pro",
name="openai/gpt-5.5-pro",
family="",
provider_id="openrouter",
cost_input=25.0,
cost_output=125.0,
)
assert (
expensive_model_warning(
"openai/gpt-5.5-pro",
provider=provider,
model_info=info,
)
is None
)
def test_skips_untrusted_provider_pricing_lookup_for_custom_provider(monkeypatch):
monkeypatch.setattr("agent.models_dev.get_model_info", lambda *_args, **_kwargs: None)
pricing_calls = []
def fake_get_pricing_entry(*_args, **_kwargs):
pricing_calls.append(_args)
return PricingEntry(
input_cost_per_million=Decimal("25"),
output_cost_per_million=Decimal("125"),
source="provider_models_api",
)
monkeypatch.setattr("agent.usage_pricing.get_pricing_entry", fake_get_pricing_entry)
warning = expensive_model_warning(
"openai/gpt-5.5-pro",
provider="custom:routerai",
base_url="https://routerai.example/v1",
)
assert warning is None
assert pricing_calls == []
def test_warns_when_pricing_entry_output_price_exceeds_threshold(monkeypatch):
monkeypatch.setattr("agent.models_dev.get_model_info", lambda *_args, **_kwargs: None)
monkeypatch.setattr(
"agent.usage_pricing.get_pricing_entry",
lambda *_args, **_kwargs: PricingEntry(
input_cost_per_million=Decimal("1.00"),
output_cost_per_million=Decimal("100.01"),
source="provider_models_api",
),
)
warning = expensive_model_warning("provider/expensive-output", provider="openrouter")
assert warning is not None
assert warning.output_cost_per_million == Decimal("100.01")
assert "$100.01/M" in warning.message
def test_openai_gpt55_pro_adds_suggestion(monkeypatch):
monkeypatch.setattr("agent.models_dev.get_model_info", lambda *_args, **_kwargs: None)
monkeypatch.setattr(
"agent.usage_pricing.get_pricing_entry",
lambda *_args, **_kwargs: PricingEntry(
input_cost_per_million=Decimal("25"),
output_cost_per_million=Decimal("125"),
source="provider_models_api",
),
)
warning = expensive_model_warning("openai/gpt-5.5-pro", provider="openrouter")
assert warning is not None
assert "did you mean to select openai/gpt-5.5?" in warning.message
def test_openai_gpt55_pro_warns_for_nous_portal_pricing(monkeypatch):