mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Fix cheapest model catalog selection
This commit is contained in:
+42
-60
@@ -123,35 +123,26 @@ def create_model_mappings(
|
||||
|
||||
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
|
||||
unique_models: dict[str, "Model"] = {}
|
||||
unique_model_keys: dict[str, str] = {}
|
||||
seen_model_provider: set[tuple[str, str]] = set()
|
||||
|
||||
# Providers sharing a URL may use different credentials and expose different
|
||||
# deployments. Keep them all; candidates for the same model are ranked by
|
||||
# their fee-adjusted pricing below.
|
||||
providers_by_db_id: dict[int, "BaseUpstreamProvider"] = {}
|
||||
for upstream in upstreams:
|
||||
db_id = getattr(upstream, "db_id", None)
|
||||
if isinstance(db_id, int):
|
||||
providers_by_db_id[db_id] = upstream
|
||||
|
||||
# Group upstreams by URL and keep only the one with the lowest fee for each URL
|
||||
upstreams_by_url: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
for upstream in upstreams:
|
||||
url = getattr(upstream, "base_url", "")
|
||||
if url not in upstreams_by_url:
|
||||
upstreams_by_url[url] = []
|
||||
upstreams_by_url[url].append(upstream)
|
||||
|
||||
filtered_upstreams: list["BaseUpstreamProvider"] = []
|
||||
for providers in upstreams_by_url.values():
|
||||
best_provider = min(providers, key=lambda p: p.provider_fee)
|
||||
filtered_upstreams.append(best_provider)
|
||||
|
||||
# Separate OpenRouter from other providers
|
||||
openrouter: "BaseUpstreamProvider" | None = None
|
||||
openrouter_upstreams: list["BaseUpstreamProvider"] = []
|
||||
other_upstreams: list["BaseUpstreamProvider"] = []
|
||||
|
||||
for upstream in filtered_upstreams:
|
||||
for upstream in upstreams:
|
||||
base_url = getattr(upstream, "base_url", "")
|
||||
if base_url == "https://openrouter.ai/api/v1":
|
||||
openrouter = upstream
|
||||
openrouter_upstreams.append(upstream)
|
||||
else:
|
||||
other_upstreams.append(upstream)
|
||||
|
||||
@@ -191,9 +182,13 @@ def create_model_mappings(
|
||||
return
|
||||
alias_candidates.append((model, provider))
|
||||
|
||||
def process_provider_models(
|
||||
upstream: "BaseUpstreamProvider", is_openrouter: bool = False
|
||||
) -> None:
|
||||
def record_unique_model_key(model: "Model") -> None:
|
||||
"""Record one case-insensitive public ID and its display spelling."""
|
||||
base_id = get_base_model_id(model.id)
|
||||
public_id = get_effective_forwarded_model_id(model) or base_id
|
||||
unique_model_keys.setdefault(public_id.lower(), public_id)
|
||||
|
||||
def process_provider_models(upstream: "BaseUpstreamProvider") -> None:
|
||||
"""Process all models from a given provider."""
|
||||
upstream_prefix = getattr(upstream, "upstream_name", None)
|
||||
provider_key = get_provider_identity(upstream)
|
||||
@@ -219,19 +214,7 @@ def create_model_mappings(
|
||||
else:
|
||||
model_to_use = model
|
||||
|
||||
# Add to unique models
|
||||
base_id = get_base_model_id(model_to_use.id)
|
||||
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
||||
unique_key = forwarded_model_id or base_id
|
||||
if not is_openrouter or unique_key not in unique_models:
|
||||
unique_model = model_to_use.copy(
|
||||
update={
|
||||
"id": base_id,
|
||||
"upstream_provider_id": upstream.provider_type,
|
||||
"forwarded_model_id": forwarded_model_id,
|
||||
}
|
||||
)
|
||||
unique_models[unique_key] = unique_model
|
||||
|
||||
# Get all aliases for this model
|
||||
aliases = resolve_model_alias(
|
||||
@@ -253,15 +236,16 @@ def create_model_mappings(
|
||||
# Try to set each alias
|
||||
for alias in aliases:
|
||||
_add_candidate(alias, model_to_use, upstream)
|
||||
record_unique_model_key(model_to_use)
|
||||
seen_model_provider.add((model_to_use.id.lower(), provider_key))
|
||||
|
||||
# Process non-OpenRouter providers first
|
||||
for upstream in other_upstreams:
|
||||
process_provider_models(upstream, is_openrouter=False)
|
||||
process_provider_models(upstream)
|
||||
|
||||
# Process OpenRouter last
|
||||
if openrouter:
|
||||
process_provider_models(openrouter, is_openrouter=True)
|
||||
for upstream in openrouter_upstreams:
|
||||
process_provider_models(upstream)
|
||||
|
||||
# Include enabled DB overrides even when provider discovery misses models.
|
||||
# This is important for deployment-based providers like Azure.
|
||||
@@ -297,22 +281,7 @@ def create_model_mappings(
|
||||
if not model_to_use.enabled:
|
||||
continue
|
||||
|
||||
base_id = get_base_model_id(model_to_use.id)
|
||||
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
||||
unique_key = forwarded_model_id or base_id
|
||||
is_openrouter = (
|
||||
getattr(upstream_for_override, "base_url", "")
|
||||
== "https://openrouter.ai/api/v1"
|
||||
)
|
||||
if not is_openrouter or unique_key not in unique_models:
|
||||
unique_model = model_to_use.copy(
|
||||
update={
|
||||
"id": base_id,
|
||||
"upstream_provider_id": upstream_for_override.provider_type,
|
||||
"forwarded_model_id": forwarded_model_id,
|
||||
}
|
||||
)
|
||||
unique_models[unique_key] = unique_model
|
||||
|
||||
try:
|
||||
aliases = resolve_model_alias(
|
||||
@@ -344,6 +313,7 @@ def create_model_mappings(
|
||||
|
||||
for alias in aliases:
|
||||
_add_candidate(alias, model_to_use, upstream_for_override)
|
||||
record_unique_model_key(model_to_use)
|
||||
seen_model_provider.add(dedupe_key)
|
||||
|
||||
# Sort candidates and build final maps
|
||||
@@ -353,20 +323,16 @@ def create_model_mappings(
|
||||
def alias_priority(model: "Model", alias: str) -> int:
|
||||
"""Rank how strong the mapping of alias->model is.
|
||||
|
||||
forwarded_model_id is the most specific identifier (set per-provider
|
||||
instance), so a match there should beat a model_id match. This way,
|
||||
when multiple providers have the same model_id but different
|
||||
forwarded_model_ids, the one whose forwarded_model_id equals the
|
||||
requested alias wins.
|
||||
An exact model ID is authoritative and must be cost-ranked against the
|
||||
other exact matches before considering forwarded aliases. This keeps a
|
||||
provider-specific forwarded ID from shadowing a directly available,
|
||||
cheaper model with the requested ID.
|
||||
"""
|
||||
forwarded_model_id = get_effective_forwarded_model_id(model)
|
||||
if forwarded_model_id and forwarded_model_id.lower() == alias:
|
||||
if model.id and model.id.lower() == alias:
|
||||
return 5
|
||||
|
||||
if (
|
||||
model.id
|
||||
and model.id.lower() == alias
|
||||
):
|
||||
forwarded_model_id = get_effective_forwarded_model_id(model)
|
||||
if forwarded_model_id and forwarded_model_id.lower() == alias:
|
||||
return 4
|
||||
|
||||
model_base = get_base_model_id(model.id)
|
||||
@@ -395,6 +361,22 @@ def create_model_mappings(
|
||||
model_instances[alias] = best_model
|
||||
provider_map[alias] = list(items)
|
||||
|
||||
# The catalog must advertise the same provider-specific model that routing
|
||||
# selects for the public ID. Normally this is the cheapest candidate; alias
|
||||
# priority intentionally wins forwarded-ID collisions.
|
||||
for unique_key, advertised_id in unique_model_keys.items():
|
||||
ranked_candidates = provider_map.get(unique_key)
|
||||
if not ranked_candidates:
|
||||
continue
|
||||
best_model, best_provider = ranked_candidates[0]
|
||||
unique_models[unique_key] = best_model.copy(
|
||||
update={
|
||||
"id": advertised_id,
|
||||
"upstream_provider_id": best_provider.provider_type,
|
||||
"forwarded_model_id": get_effective_forwarded_model_id(best_model),
|
||||
}
|
||||
)
|
||||
|
||||
# Log provider distribution (using top provider for stats)
|
||||
provider_counts: dict[str, int] = {}
|
||||
for candidate_list in provider_map.values():
|
||||
|
||||
@@ -618,10 +618,8 @@ async def model_paths_for_model(model_id: str) -> dict:
|
||||
|
||||
result = await get_paths_for_model(model_id)
|
||||
if not result["data"]:
|
||||
advertised_ids = {
|
||||
model.forwarded_model_id or model.id for model in get_unique_models()
|
||||
}
|
||||
if model_id not in advertised_ids:
|
||||
advertised_ids = {model.id.lower() for model in get_unique_models()}
|
||||
if model_id.lower() not in advertised_ids:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
return result
|
||||
|
||||
@@ -637,8 +635,5 @@ async def models(session: AsyncSession = Depends(get_session)) -> dict:
|
||||
items = get_unique_models()
|
||||
data = []
|
||||
for model in items:
|
||||
m = model.dict()
|
||||
if model.forwarded_model_id:
|
||||
m["id"] = model.forwarded_model_id
|
||||
data.append(m)
|
||||
data.append(model.dict())
|
||||
return {"data": data}
|
||||
|
||||
@@ -106,14 +106,8 @@ async def dual_provider_maps(
|
||||
yield cheap, expensive
|
||||
|
||||
|
||||
def _upstream_response(request: httpx.Request) -> httpx.Response:
|
||||
"""502 from the cheap (winning) provider; a served completion elsewhere."""
|
||||
if request.url.host == "cheap.example.com":
|
||||
return httpx.Response(
|
||||
502,
|
||||
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
def _successful_upstream_response() -> httpx.Response:
|
||||
"""Return a successful completion with deterministic token usage."""
|
||||
body = {
|
||||
"id": "chatcmpl-served",
|
||||
"object": "chat.completion",
|
||||
@@ -139,6 +133,17 @@ def _upstream_response(request: httpx.Request) -> httpx.Response:
|
||||
)
|
||||
|
||||
|
||||
def _upstream_response(request: httpx.Request) -> httpx.Response:
|
||||
"""502 from the cheap (winning) provider; a served completion elsewhere."""
|
||||
if request.url.host == "cheap.example.com":
|
||||
return httpx.Response(
|
||||
502,
|
||||
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
return _successful_upstream_response()
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_failover_serve_billed_at_serving_providers_rate(
|
||||
@@ -235,6 +240,44 @@ async def same_id_provider_maps(
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_id_cheapest_provider_serves_and_bills_its_own_price(
|
||||
authenticated_client: AsyncClient,
|
||||
same_id_provider_maps: None,
|
||||
) -> None:
|
||||
"""Exact same-ID candidates route cheapest and bill that provider's model."""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _successful_upstream_response()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [request.url.host for request in sent_requests] == ["cheap.example.com"]
|
||||
assert response.json()["cost"]["total_msats"] == 2_000
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_id_failover_settles_at_serving_price(
|
||||
|
||||
@@ -2,9 +2,12 @@
|
||||
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
# Set required env vars before importing
|
||||
os.environ["UPSTREAM_BASE_URL"] = "http://test"
|
||||
@@ -15,7 +18,13 @@ from routstr.algorithm import ( # noqa: E402
|
||||
create_model_mappings,
|
||||
get_provider_penalty,
|
||||
)
|
||||
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
|
||||
from routstr.core.db import get_session # noqa: E402
|
||||
from routstr.payment.models import ( # noqa: E402
|
||||
Architecture,
|
||||
Model,
|
||||
Pricing,
|
||||
models_router,
|
||||
)
|
||||
|
||||
|
||||
def create_test_model(
|
||||
@@ -56,6 +65,7 @@ def create_test_provider(
|
||||
db_id: int | None = None,
|
||||
models: list[Model] | None = None,
|
||||
upstream_name: str | None = None,
|
||||
provider_fee: float = 1.0,
|
||||
) -> Mock:
|
||||
"""Helper to create a test provider mock."""
|
||||
provider = Mock()
|
||||
@@ -63,6 +73,7 @@ def create_test_provider(
|
||||
provider.base_url = base_url
|
||||
provider.db_id = db_id
|
||||
provider.upstream_name = upstream_name or name
|
||||
provider.provider_fee = provider_fee
|
||||
provider.get_cached_models.return_value = models or []
|
||||
return provider
|
||||
|
||||
@@ -115,6 +126,537 @@ def test_get_provider_penalty_openrouter() -> None:
|
||||
assert penalty == 1.001
|
||||
|
||||
|
||||
def test_create_model_mappings_advertises_cheapest_custom_provider_regardless_of_order() -> (
|
||||
None
|
||||
):
|
||||
"""The public model catalog must use the same cheapest custom-provider model."""
|
||||
cheap_model = create_test_model(
|
||||
"shared-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
expensive_model = create_test_model(
|
||||
"shared-model", prompt_price=0.1, completion_price=0.1
|
||||
)
|
||||
cheap_provider = create_test_provider(
|
||||
"custom-cheap",
|
||||
"https://cheap.example/v1",
|
||||
db_id=1,
|
||||
models=[cheap_model],
|
||||
)
|
||||
expensive_provider = create_test_provider(
|
||||
"custom-expensive",
|
||||
"https://expensive.example/v1",
|
||||
db_id=2,
|
||||
models=[expensive_model],
|
||||
)
|
||||
|
||||
provider_orders: list[list[Any]] = [
|
||||
[cheap_provider, expensive_provider],
|
||||
[expensive_provider, cheap_provider],
|
||||
]
|
||||
for providers in provider_orders:
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=providers,
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].pricing.prompt == 0.001
|
||||
assert provider_map["shared-model"][0] == (cheap_model, cheap_provider)
|
||||
|
||||
|
||||
def test_create_model_mappings_advertises_cheaper_openrouter_model() -> None:
|
||||
"""OpenRouter may win when its adjusted price beats every custom provider."""
|
||||
custom_model = create_test_model(
|
||||
"shared-model", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
openrouter_model = create_test_model(
|
||||
"shared-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
custom = create_test_provider(
|
||||
"custom", "https://custom.example/v1", db_id=1, models=[custom_model]
|
||||
)
|
||||
openrouter = create_test_provider(
|
||||
"openrouter",
|
||||
"https://openrouter.ai/api/v1",
|
||||
db_id=2,
|
||||
models=[openrouter_model],
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[openrouter, custom],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].pricing.prompt == 0.001
|
||||
assert provider_map["shared-model"][0] == (openrouter_model, openrouter)
|
||||
|
||||
|
||||
def test_create_model_mappings_ranks_multiple_openrouter_credentials() -> None:
|
||||
"""OpenRouter accounts sharing a URL still compete by effective model cost."""
|
||||
cheap_model = create_test_model(
|
||||
"shared-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
expensive_model = create_test_model(
|
||||
"shared-model", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
cheap = create_test_provider(
|
||||
"openrouter-cheap",
|
||||
"https://openrouter.ai/api/v1",
|
||||
db_id=1,
|
||||
models=[cheap_model],
|
||||
)
|
||||
expensive = create_test_provider(
|
||||
"openrouter-expensive",
|
||||
"https://openrouter.ai/api/v1",
|
||||
db_id=2,
|
||||
models=[expensive_model],
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[cheap, expensive],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert provider_map["shared-model"][0] == (cheap_model, cheap)
|
||||
assert unique_models["shared-model"].upstream_provider_id == "openrouter-cheap"
|
||||
|
||||
|
||||
def test_create_model_mappings_uses_openrouter_penalty_for_catalog_ties() -> None:
|
||||
"""Equal raw prices prefer a custom provider in routing and the catalog."""
|
||||
custom_model = create_test_model("shared-model")
|
||||
openrouter_model = create_test_model("shared-model")
|
||||
custom = create_test_provider(
|
||||
"custom", "https://custom.example/v1", db_id=1, models=[custom_model]
|
||||
)
|
||||
openrouter = create_test_provider(
|
||||
"openrouter",
|
||||
"https://openrouter.ai/api/v1",
|
||||
db_id=2,
|
||||
models=[openrouter_model],
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[openrouter, custom],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].upstream_provider_id == "custom"
|
||||
assert provider_map["shared-model"][0] == (custom_model, custom)
|
||||
|
||||
|
||||
def test_create_model_mappings_applies_custom_provider_fees_before_advertising(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Custom-provider DB fees participate in cheapest-model selection."""
|
||||
providers: list[Any] = [
|
||||
create_test_provider(
|
||||
"high-fee",
|
||||
"https://high-fee.example/v1",
|
||||
db_id=1,
|
||||
models=[create_test_model("shared-model")],
|
||||
provider_fee=1.20,
|
||||
),
|
||||
create_test_provider(
|
||||
"low-fee",
|
||||
"https://low-fee.example/v1",
|
||||
db_id=2,
|
||||
models=[create_test_model("shared-model")],
|
||||
provider_fee=1.01,
|
||||
),
|
||||
]
|
||||
rows = {
|
||||
1: SimpleNamespace(id="shared-model", upstream_provider_id=1, enabled=True),
|
||||
2: SimpleNamespace(id="shared-model", upstream_provider_id=2, enabled=True),
|
||||
}
|
||||
|
||||
def fake_row_to_model(row, *, apply_provider_fee, provider_fee) -> Model: # type: ignore[no-untyped-def]
|
||||
assert apply_provider_fee is True
|
||||
return create_test_model(
|
||||
row.id,
|
||||
prompt_price=0.001 * provider_fee,
|
||||
completion_price=0.002 * provider_fee,
|
||||
)
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=providers,
|
||||
overrides_by_key={
|
||||
("shared-model", provider_id): (
|
||||
row,
|
||||
providers[provider_id - 1].provider_fee,
|
||||
)
|
||||
for provider_id, row in rows.items()
|
||||
},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].upstream_provider_id == "low-fee"
|
||||
assert provider_map["shared-model"][0][1] is providers[1]
|
||||
|
||||
|
||||
def test_create_model_mappings_compares_missing_custom_override_with_openrouter(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""A custom model created in the DB can beat a close OpenRouter candidate."""
|
||||
custom = create_test_provider(
|
||||
"custom", "https://custom.example/v1", db_id=1, models=[]
|
||||
)
|
||||
openrouter_model = create_test_model(
|
||||
"shared-model", prompt_price=0.0009995, completion_price=0.0009995
|
||||
)
|
||||
openrouter = create_test_provider(
|
||||
"openrouter",
|
||||
"https://openrouter.ai/api/v1",
|
||||
db_id=2,
|
||||
models=[openrouter_model],
|
||||
)
|
||||
custom_override = create_test_model(
|
||||
"shared-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
override_row = SimpleNamespace(
|
||||
id="shared-model", upstream_provider_id=1, enabled=True
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"routstr.payment.models._row_to_model", lambda *args, **kwargs: custom_override
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[openrouter, custom],
|
||||
overrides_by_key={("shared-model", 1): (override_row, 1.0)},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].upstream_provider_id == "custom"
|
||||
assert provider_map["shared-model"][0] == (custom_override, custom)
|
||||
|
||||
|
||||
def test_create_model_mappings_excludes_disabled_cheapest_custom_model() -> None:
|
||||
"""A disabled provider-scoped model cannot become the advertised cheapest."""
|
||||
disabled_cheap = create_test_model(
|
||||
"shared-model", prompt_price=0.0001, completion_price=0.0001
|
||||
)
|
||||
enabled_expensive = create_test_model(
|
||||
"shared-model", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
cheap_provider = create_test_provider(
|
||||
"disabled-cheap",
|
||||
"https://cheap.example/v1",
|
||||
db_id=1,
|
||||
models=[disabled_cheap],
|
||||
)
|
||||
enabled_provider = create_test_provider(
|
||||
"enabled",
|
||||
"https://enabled.example/v1",
|
||||
db_id=2,
|
||||
models=[enabled_expensive],
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[cheap_provider, enabled_provider],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys={("shared-model", 1)},
|
||||
)
|
||||
|
||||
assert unique_models["shared-model"].upstream_provider_id == "enabled"
|
||||
assert provider_map["shared-model"] == [(enabled_expensive, enabled_provider)]
|
||||
|
||||
|
||||
def test_create_model_mappings_same_url_ranks_shared_model_and_keeps_unique_override(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Same-URL providers compete by cost without losing provider-only models."""
|
||||
shared_url = "https://custom.example/v1"
|
||||
high_fee = create_test_provider(
|
||||
"high-fee",
|
||||
shared_url,
|
||||
db_id=1,
|
||||
models=[create_test_model("shared-model", prompt_price=0.0012)],
|
||||
provider_fee=1.20,
|
||||
)
|
||||
low_fee = create_test_provider(
|
||||
"low-fee",
|
||||
shared_url,
|
||||
db_id=2,
|
||||
models=[create_test_model("shared-model", prompt_price=0.00101)],
|
||||
provider_fee=1.01,
|
||||
)
|
||||
unique_override = SimpleNamespace(
|
||||
id="unique-deployment", upstream_provider_id=1, enabled=True
|
||||
)
|
||||
override_model = create_test_model("unique-deployment", prompt_price=0.002)
|
||||
monkeypatch.setattr(
|
||||
"routstr.payment.models._row_to_model", Mock(return_value=override_model)
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[high_fee, low_fee],
|
||||
overrides_by_key={("unique-deployment", 1): (unique_override, 1.20)},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert provider_map["shared-model"][0][1] is low_fee
|
||||
assert provider_map["unique-deployment"] == [(override_model, high_fee)]
|
||||
assert set(unique_models) == {"shared-model", "unique-deployment"}
|
||||
|
||||
|
||||
def test_create_model_mappings_chooses_cheapest_shared_forwarded_id() -> None:
|
||||
"""Custom deployment names sharing a public ID advertise the cheapest deployment."""
|
||||
expensive = create_test_model(
|
||||
"deployment-expensive", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
expensive.forwarded_model_id = "public-model"
|
||||
cheap = create_test_model(
|
||||
"deployment-cheap", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
cheap.forwarded_model_id = "public-model"
|
||||
expensive_provider = create_test_provider(
|
||||
"expensive",
|
||||
"https://expensive.example/v1",
|
||||
db_id=1,
|
||||
models=[expensive],
|
||||
)
|
||||
cheap_provider = create_test_provider(
|
||||
"cheap", "https://cheap.example/v1", db_id=2, models=[cheap]
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[cheap_provider, expensive_provider],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert unique_models["public-model"].pricing.prompt == 0.001
|
||||
assert provider_map["public-model"][0] == (cheap, cheap_provider)
|
||||
|
||||
|
||||
def test_create_model_mappings_exact_model_id_beats_forwarded_id_collision(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""An exact model ID uses its cheapest provider before forwarded aliases."""
|
||||
direct_expensive = create_test_model(
|
||||
"public-model", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
direct_cheapest = create_test_model(
|
||||
"public-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
forwarded = create_test_model(
|
||||
"deployment-name", prompt_price=0.0001, completion_price=0.0001
|
||||
)
|
||||
forwarded.forwarded_model_id = "public-model"
|
||||
direct_expensive_provider = create_test_provider(
|
||||
"direct-expensive",
|
||||
"https://direct-expensive.example/v1",
|
||||
db_id=1,
|
||||
models=[direct_expensive],
|
||||
)
|
||||
direct_cheapest_provider = create_test_provider(
|
||||
"direct-cheapest",
|
||||
"https://direct-cheapest.example/v1",
|
||||
db_id=2,
|
||||
models=[direct_cheapest],
|
||||
)
|
||||
forwarded_provider = create_test_provider(
|
||||
"forwarded",
|
||||
"https://forwarded.example/v1",
|
||||
db_id=3,
|
||||
models=[forwarded],
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[
|
||||
direct_expensive_provider,
|
||||
forwarded_provider,
|
||||
direct_cheapest_provider,
|
||||
],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert provider_map["public-model"][0] == (
|
||||
direct_cheapest,
|
||||
direct_cheapest_provider,
|
||||
)
|
||||
assert unique_models["public-model"].pricing.prompt == 0.001
|
||||
assert unique_models["public-model"].upstream_provider_id == "direct-cheapest"
|
||||
|
||||
import routstr.proxy as proxy
|
||||
|
||||
monkeypatch.setattr(proxy, "_unique_models", unique_models)
|
||||
app = FastAPI()
|
||||
app.include_router(models_router)
|
||||
app.dependency_overrides[get_session] = lambda: None
|
||||
response = TestClient(app).get("/v1/models")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["data"][0]["id"] == "public-model"
|
||||
assert response.json()["data"][0]["pricing"]["prompt"] == 0.001
|
||||
|
||||
|
||||
def test_models_endpoint_preserves_catalog_id_when_winner_forwards_elsewhere(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Each catalog row keeps its requested ID while using its routing winner."""
|
||||
base_alias = create_test_model(
|
||||
"vendor/foo", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
redirected_exact = create_test_model("foo", prompt_price=0.1, completion_price=0.1)
|
||||
redirected_exact.forwarded_model_id = "bar"
|
||||
base_provider = create_test_provider(
|
||||
"base", "https://base.example/v1", db_id=1, models=[base_alias]
|
||||
)
|
||||
redirect_provider = create_test_provider(
|
||||
"redirect", "https://redirect.example/v1", db_id=2, models=[redirected_exact]
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[base_provider, redirect_provider],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert provider_map["foo"][0] == (redirected_exact, redirect_provider)
|
||||
assert unique_models["foo"].id == "foo"
|
||||
assert unique_models["foo"].upstream_provider_id == "redirect"
|
||||
|
||||
import routstr.proxy as proxy
|
||||
|
||||
monkeypatch.setattr(proxy, "_unique_models", unique_models)
|
||||
app = FastAPI()
|
||||
app.include_router(models_router)
|
||||
app.dependency_overrides[get_session] = lambda: None
|
||||
response = TestClient(app).get("/v1/models")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert {model["id"] for model in response.json()["data"]} == {"foo", "bar"}
|
||||
|
||||
with patch(
|
||||
"routstr.upstream.model_paths.get_paths_for_model",
|
||||
new=AsyncMock(return_value={"data": []}),
|
||||
):
|
||||
path_response = TestClient(app).get(
|
||||
"/v1/models/paths/model", params={"model_id": "foo"}
|
||||
)
|
||||
assert path_response.status_code == 200
|
||||
|
||||
|
||||
def test_models_endpoint_dedupes_case_insensitive_model_ids(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Case-insensitive routing aliases produce one catalog row, not duplicates."""
|
||||
upper = create_test_model("GPT-4o", prompt_price=0.001, completion_price=0.001)
|
||||
lower = create_test_model("gpt-4o", prompt_price=0.01, completion_price=0.01)
|
||||
upper_provider = create_test_provider(
|
||||
"upper", "https://upper.example/v1", db_id=1, models=[upper]
|
||||
)
|
||||
lower_provider = create_test_provider(
|
||||
"lower", "https://lower.example/v1", db_id=2, models=[lower]
|
||||
)
|
||||
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[upper_provider, lower_provider],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert list(unique_models) == ["gpt-4o"]
|
||||
assert provider_map["gpt-4o"][0] == (upper, upper_provider)
|
||||
|
||||
import routstr.proxy as proxy
|
||||
|
||||
monkeypatch.setattr(proxy, "_unique_models", unique_models)
|
||||
app = FastAPI()
|
||||
app.include_router(models_router)
|
||||
app.dependency_overrides[get_session] = lambda: None
|
||||
response = TestClient(app).get("/v1/models")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["id"] for model in response.json()["data"]] == ["GPT-4o"]
|
||||
|
||||
with patch(
|
||||
"routstr.upstream.model_paths.get_paths_for_model",
|
||||
new=AsyncMock(return_value={"data": []}),
|
||||
):
|
||||
path_response = TestClient(app).get(
|
||||
"/v1/models/paths/model", params={"model_id": "gpt-4o"}
|
||||
)
|
||||
assert path_response.status_code == 200
|
||||
|
||||
|
||||
def test_create_model_mappings_equal_custom_prices_remain_cheapest() -> None:
|
||||
"""Tied custom providers advertise one of the equally cheapest candidates."""
|
||||
first_model = create_test_model("shared-model")
|
||||
second_model = create_test_model("shared-model")
|
||||
first = create_test_provider(
|
||||
"first", "https://first.example/v1", db_id=1, models=[first_model]
|
||||
)
|
||||
second = create_test_provider(
|
||||
"second", "https://second.example/v1", db_id=2, models=[second_model]
|
||||
)
|
||||
|
||||
provider_orders: list[list[Any]] = [[first, second], [second, first]]
|
||||
for providers in provider_orders:
|
||||
_, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=providers,
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert calculate_model_cost_score(unique_models["shared-model"]) == min(
|
||||
calculate_model_cost_score(first_model),
|
||||
calculate_model_cost_score(second_model),
|
||||
)
|
||||
assert (
|
||||
unique_models["shared-model"].upstream_provider_id
|
||||
== provider_map["shared-model"][0][1].provider_type
|
||||
)
|
||||
|
||||
|
||||
def test_models_endpoint_returns_cheapest_custom_provider_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""GET /v1/models exposes the cheapest model chosen from custom providers."""
|
||||
cheap = create_test_model(
|
||||
"shared-model", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
expensive = create_test_model(
|
||||
"shared-model", prompt_price=0.1, completion_price=0.1
|
||||
)
|
||||
providers: list[Any] = [
|
||||
create_test_provider(
|
||||
"cheap", "https://cheap.example/v1", db_id=1, models=[cheap]
|
||||
),
|
||||
create_test_provider(
|
||||
"expensive",
|
||||
"https://expensive.example/v1",
|
||||
db_id=2,
|
||||
models=[expensive],
|
||||
),
|
||||
]
|
||||
_, _, unique_models = create_model_mappings(
|
||||
upstreams=providers,
|
||||
overrides_by_key={},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
import routstr.proxy as proxy
|
||||
|
||||
monkeypatch.setattr(proxy, "_unique_models", unique_models)
|
||||
app = FastAPI()
|
||||
app.include_router(models_router)
|
||||
app.dependency_overrides[get_session] = lambda: None
|
||||
|
||||
response = TestClient(app).get("/v1/models")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert len(response.json()["data"]) == 1
|
||||
assert response.json()["data"][0]["pricing"]["prompt"] == 0.001
|
||||
assert response.json()["data"][0]["upstream_provider_id"] == "cheap"
|
||||
|
||||
|
||||
def test_create_model_mappings_includes_db_override_for_missing_cached_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
@@ -133,7 +675,9 @@ def test_create_model_mappings_includes_db_override_for_missing_cached_model(
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=7, enabled=True)
|
||||
override_row = SimpleNamespace(
|
||||
id="azure/gpt-4o", upstream_provider_id=7, enabled=True
|
||||
)
|
||||
|
||||
model_instances, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[provider],
|
||||
@@ -178,7 +722,9 @@ def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=2, enabled=True)
|
||||
override_row = SimpleNamespace(
|
||||
id="azure/gpt-4o", upstream_provider_id=2, enabled=True
|
||||
)
|
||||
|
||||
_, provider_map, _ = create_model_mappings(
|
||||
upstreams=[provider_a, provider_b],
|
||||
@@ -246,9 +792,7 @@ def test_create_model_mappings_does_not_split_self_alias_from_base_identity(
|
||||
"https://provider-b.example/v1",
|
||||
db_id=2,
|
||||
models=[
|
||||
create_test_model(
|
||||
model_id, prompt_price=0.0001, completion_price=0.0001
|
||||
)
|
||||
create_test_model(model_id, prompt_price=0.0001, completion_price=0.0001)
|
||||
],
|
||||
)
|
||||
|
||||
@@ -304,8 +848,9 @@ def test_create_model_mappings_preserves_case_only_forwarded_alias(
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert list(unique_models) == [case_only_alias]
|
||||
assert unique_models[case_only_alias].forwarded_model_id == case_only_alias
|
||||
normalized_alias = case_only_alias.lower()
|
||||
assert list(unique_models) == [normalized_alias]
|
||||
assert unique_models[normalized_alias].forwarded_model_id == case_only_alias
|
||||
assert len(provider_map[case_only_alias.lower()]) == 1
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user