fix(model): prefer the current provider's alias on implicit switches too

Only the --provider path handed user_providers/custom_providers to
resolve_alias, so the provider-identity comparison behind the
current-provider preference could not see legacy custom_providers
entries on the implicit path (/model <id>) or the authenticated-provider
fallback. On custom:corp-llm, /model shared-model picked another
provider's alias for the same model id and switched to its base_url.

Pass both provider maps on every resolve_alias call site. The two
existing resolve_alias stubs in test_ollama_cloud_auth.py accept the
extra positional args; their assertions are unchanged.
This commit is contained in:
teknium1
2026-09-23 16:31:34 -07:00
committed by Teknium
parent 79d012bd25
commit ffc44bbaf3
2 changed files with 43 additions and 6 deletions

View File

@@ -855,12 +855,14 @@ def get_authenticated_provider_slugs(
def _resolve_alias_fallback(
raw_input: str, authenticated_providers: list[str] = ()) -> Optional[tuple[str, str, str]]:
raw_input: str, authenticated_providers: list[str] = (), user_providers: Optional[dict] = None,
custom_providers: Optional[list] = None) -> Optional[tuple[str, str, str]]:
"""Resolve an alias on the user's authenticated providers (``("openrouter", "nous")`` when none given).
AmbiguousAliasError propagates: the alias exists on this provider, the user just has to
choose — trying the next provider would silently switch them somewhere they didn't ask for."""
results = (resolve_alias(raw_input, p) for p in authenticated_providers or ("openrouter", "nous"))
results = (resolve_alias(raw_input, p, user_providers, custom_providers)
for p in authenticated_providers or ("openrouter", "nous"))
return next((r for r in results if r is not None), None)
@@ -1267,7 +1269,7 @@ def _route_alias_fallback(st: _Switch, key: str) -> Optional[ModelSwitchResult]:
current_provider=st.current_provider, user_providers=st.user_providers, custom_providers=st.custom_providers,
)
try:
fallback_result = _resolve_alias_fallback(st.raw_input, authed)
fallback_result = _resolve_alias_fallback(st.raw_input, authed, st.user_providers, st.custom_providers)
except AmbiguousAliasError as err:
return st.fail(_ambiguous_alias_message(err))
if fallback_result is None:
@@ -1359,7 +1361,7 @@ def _route_from_model_input(st: _Switch) -> Optional[ModelSwitchResult]:
st.target_provider, st.new_model, st.resolved_alias = "moa", moa_match, ""
else:
try:
alias_result = resolve_alias(raw_input, current_provider)
alias_result = resolve_alias(raw_input, current_provider, st.user_providers, st.custom_providers)
except AmbiguousAliasError as err:
return st.fail(_ambiguous_alias_message(err))
if alias_result is not None:

View File

@@ -243,7 +243,7 @@ class TestSwitchModelDirectAliasOverride:
monkeypatch.setattr(ms, "DIRECT_ALIASES", test_aliases)
monkeypatch.setattr(ms, "resolve_alias",
lambda raw, prov: ("custom", "qwen3.5:397b", "qwen"))
lambda raw, prov, *_: ("custom", "qwen3.5:397b", "qwen"))
monkeypatch.setattr(
"hermes_cli.runtime_provider.resolve_runtime_provider",
@@ -270,7 +270,7 @@ class TestSwitchModelDirectAliasOverride:
}
monkeypatch.setattr(ms, "DIRECT_ALIASES", test_aliases)
monkeypatch.setattr(ms, "resolve_alias",
lambda raw, prov: ("custom", "local-model", "local"))
lambda raw, prov, *_: ("custom", "local-model", "local"))
monkeypatch.setattr(
"hermes_cli.runtime_provider.resolve_runtime_provider",
lambda **kwargs: {"api_key": "", "base_url": "", "api_mode": "openai_compat", "provider": "custom"},
@@ -363,3 +363,38 @@ class TestSwitchModelDirectAliasOverride:
assert result.resolved_via_alias == "corp-alias"
assert result.base_url == "https://corp.example.com/v2"
assert result.api_key == "sk-corp-alias"
def test_implicit_switch_prefers_alias_of_current_legacy_custom_provider(self, monkeypatch):
"""Without --provider the current provider still owns a shared model id: on
``custom:corp-llm`` the alias naming ``corp-llm`` wins over another provider's alias."""
import os
from pathlib import Path
import hermes_cli.model_switch as ms
from hermes_cli.config import load_config
from hermes_cli.model_switch import DirectAlias
(Path(os.environ["HERMES_HOME"]) / "config.yaml").write_text(
"model:\n provider: custom:corp-llm\n default: old-model\n"
"providers:\n provider-a:\n base_url: https://api-a.example.com/v1\n"
"custom_providers:\n - name: corp-llm\n"
" base_url: https://corp.example.com/v1\n api_key: sk-corp\n")
monkeypatch.setattr(ms, "DIRECT_ALIASES", {
"a-alias": DirectAlias("shared-model", "provider-a", "https://alias-host.example.com/v1",
api_key="sk-alias-host"),
"corp-alias": DirectAlias("shared-model", "corp-llm", "https://corp.example.com/v2",
api_key="sk-corp-alias"),
})
monkeypatch.setattr("hermes_cli.models_validate.validate_requested_model",
lambda *a, **kw: {"accepted": True, "persist": True, "recognized": True, "message": None})
cfg = load_config()
result = ms.switch_model(
"shared-model", "custom:corp-llm", "old-model",
current_base_url="https://corp.example.com/v1", current_api_key="sk-corp",
user_providers=cfg["providers"], custom_providers=cfg.get("custom_providers"))
assert result.success, result.error_message
assert result.resolved_via_alias == "corp-alias"
assert result.base_url == "https://corp.example.com/v2"
assert result.api_key == "sk-corp-alias"