diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 7fba4063..777e5110 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -86,8 +86,8 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: def create_model_mappings( upstreams: list["BaseUpstreamProvider"], - overrides_by_id: dict[str, tuple], - disabled_model_ids: set[str], + overrides_by_key: dict[tuple[str, int], tuple], + disabled_model_keys: set[tuple[str, int]], ) -> tuple[ dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"] ]: @@ -107,8 +107,9 @@ def create_model_mappings( Args: upstreams: List of all upstream provider instances - overrides_by_id: Dict of model overrides from database {model_id: (ModelRow, fee)} - disabled_model_ids: Set of model IDs that should be excluded + overrides_by_key: Dict of model overrides from database + {(model_id_lower, upstream_provider_id): (ModelRow, fee)} + disabled_model_keys: Set of provider-scoped model keys that should be excluded Returns: Tuple of (model_instances, provider_map, unique_models) @@ -179,14 +180,22 @@ def create_model_mappings( """Process all models from a given provider.""" upstream_prefix = getattr(upstream, "upstream_name", None) provider_key = get_provider_identity(upstream) + upstream_db_id = getattr(upstream, "db_id", None) for model in upstream.get_cached_models(): - if not model.enabled or model.id in disabled_model_ids: + model_key = ( + (model.id.lower(), upstream_db_id) + if isinstance(upstream_db_id, int) + else None + ) + if not model.enabled or ( + model_key is not None and model_key in disabled_model_keys + ): continue - # Apply overrides if present - if model.id in overrides_by_id: - override_row, provider_fee = overrides_by_id[model.id] + # Apply overrides only for this provider's model row. + if model_key is not None and model_key in overrides_by_key: + override_row, provider_fee = overrides_by_key[model_key] model_to_use = _row_to_model( override_row, apply_provider_fee=True, provider_fee=provider_fee ) @@ -237,13 +246,10 @@ def create_model_mappings( # Include enabled DB overrides even when provider discovery misses models. # This is important for deployment-based providers like Azure. - for model_id, override_data in overrides_by_id.items(): - if model_id in disabled_model_ids: + for (model_id, upstream_provider_id), override_data in overrides_by_key.items(): + if (model_id, upstream_provider_id) in disabled_model_keys: continue override_row, provider_fee = override_data - upstream_provider_id = getattr(override_row, "upstream_provider_id", None) - if not isinstance(upstream_provider_id, int): - continue upstream_for_override = providers_by_db_id.get(upstream_provider_id) if upstream_for_override is None: diff --git a/routstr/proxy.py b/routstr/proxy.py index 8f6564e7..1bfb23e0 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -118,22 +118,23 @@ async def refresh_model_maps() -> None: result = await session.exec(query) provider_rows = result.all() - overrides_by_id: dict[str, tuple[ModelRow, float]] = {} - disabled_model_ids: set[str] = set() + overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {} + disabled_model_keys: set[tuple[str, int]] = set() for provider in provider_rows: if not provider.enabled: continue for model in provider.models: + model_key = (model.id.lower(), model.upstream_provider_id) if model.enabled: - overrides_by_id[model.id] = (model, provider.provider_fee) + overrides_by_key[model_key] = (model, provider.provider_fee) else: - disabled_model_ids.add(model.id) + disabled_model_keys.add(model_key) _model_instances, _provider_map, _unique_models = create_model_mappings( upstreams=_upstreams, - overrides_by_id=overrides_by_id, - disabled_model_ids=disabled_model_ids, + overrides_by_key=overrides_by_key, + disabled_model_keys=disabled_model_keys, ) diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index c2bc2acb..2ca3c85e 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -94,12 +94,10 @@ async def get_all_models_with_overrides( provider_result = await session.exec(select(UpstreamProviderRow)) providers_by_id = {p.id: p for p in provider_result.all()} - overrides_by_id: dict[str, tuple[ModelRow, float]] = { - row.id: ( + overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = { + (row.id.lower(), row.upstream_provider_id): ( row, - providers_by_id[row.upstream_provider_id].provider_fee - if row.upstream_provider_id in providers_by_id - else 1.01, + providers_by_id[row.upstream_provider_id].provider_fee, ) for row in override_rows if row.upstream_provider_id is not None @@ -107,17 +105,28 @@ async def get_all_models_with_overrides( and providers_by_id[row.upstream_provider_id].enabled } - all_models: dict[str, Model] = {} + all_models: dict[tuple[str, str], Model] = {} for upstream in upstreams: + upstream_db_id = getattr(upstream, "db_id", None) + provider_key = ( + f"db:{upstream_db_id}" + if isinstance(upstream_db_id, int) + else f"{getattr(upstream, 'provider_type', '')}|{getattr(upstream, 'base_url', '')}" + ) for model in upstream.get_cached_models(): - if model.id in overrides_by_id: - override_row, provider_fee = overrides_by_id[model.id] - all_models[model.id] = _row_to_model( + model_key = ( + (model.id.lower(), upstream_db_id) + if isinstance(upstream_db_id, int) + else None + ) + if model_key is not None and model_key in overrides_by_key: + override_row, provider_fee = overrides_by_key[model_key] + all_models[(model.id.lower(), provider_key)] = _row_to_model( override_row, apply_provider_fee=True, provider_fee=provider_fee ) elif model.enabled: - all_models[model.id] = model + all_models[(model.id.lower(), provider_key)] = model return list(all_models.values()) diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index 1a647528..b3b457cc 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -137,8 +137,8 @@ def test_create_model_mappings_includes_db_override_for_missing_cached_model( model_instances, provider_map, unique_models = create_model_mappings( upstreams=[provider], - overrides_by_id={"azure/gpt-4o": (override_row, 1.01)}, - disabled_model_ids=set(), + overrides_by_key={("azure/gpt-4o", 7): (override_row, 1.01)}, + disabled_model_keys=set(), ) assert "azure/gpt-4o" in model_instances @@ -182,11 +182,73 @@ def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type( _, provider_map, _ = create_model_mappings( upstreams=[provider_a, provider_b], - overrides_by_id={"azure/gpt-4o": (override_row, 1.01)}, - disabled_model_ids=set(), + overrides_by_key={("azure/gpt-4o", 2): (override_row, 1.01)}, + disabled_model_keys=set(), ) providers_for_alias = provider_map["azure/gpt-4o"] assert provider_a in providers_for_alias assert provider_b in providers_for_alias assert len(providers_for_alias) == 2 + + +def test_create_model_mappings_applies_override_only_to_matching_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Same-id overrides must not add provider-specific aliases to other providers.""" + provider_a_model = create_test_model("same-id", prompt_price=0.01) + provider_a = create_test_provider( + "provider-a", + "https://provider-a.example/v1", + db_id=1, + models=[provider_a_model], + ) + provider_b_model = create_test_model("same-id", prompt_price=0.02) + provider_b = create_test_provider( + "provider-b", + "https://provider-b.example/v1", + db_id=2, + models=[provider_b_model], + ) + + override_model = create_test_model("same-id", prompt_price=0.001) + override_model.alias_ids = ["provider-b-only"] + override_row = SimpleNamespace(id="same-id", upstream_provider_id=2, enabled=True) + + def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def] + return override_model + + monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model) + + _, provider_map, _ = create_model_mappings( + upstreams=[provider_a, provider_b], + overrides_by_key={("same-id", 2): (override_row, 1.01)}, + disabled_model_keys=set(), + ) + + assert provider_map["provider-b-only"] == [provider_b] + assert set(provider_map["same-id"]) == {provider_a, provider_b} + + +def test_create_model_mappings_disables_only_matching_provider() -> None: + """Disabled overrides are scoped to the provider row, not the shared model id.""" + provider_a = create_test_provider( + "provider-a", + "https://provider-a.example/v1", + db_id=1, + models=[create_test_model("same-id")], + ) + provider_b = create_test_provider( + "provider-b", + "https://provider-b.example/v1", + db_id=2, + models=[create_test_model("same-id")], + ) + + _, provider_map, _ = create_model_mappings( + upstreams=[provider_a, provider_b], + overrides_by_key={}, + disabled_model_keys={("same-id", 2)}, + ) + + assert provider_map["same-id"] == [provider_a]