From faaac1b262249eafcedcd468c2396a6856dc2285 Mon Sep 17 00:00:00 2001 From: Brian Krabach Date: Wed, 30 Sep 2026 06:28:47 -0700 Subject: [PATCH] Honor selected provider configuration during login and discovery --- amplifier_app_cli/commands/provider.py | 30 +- amplifier_app_cli/provider_loader.py | 88 +++--- tests/test_provider_loader_configuration.py | 334 ++++++++++++++++++++ 3 files changed, 403 insertions(+), 49 deletions(-) create mode 100644 tests/test_provider_loader_configuration.py diff --git a/amplifier_app_cli/commands/provider.py b/amplifier_app_cli/commands/provider.py index 8d38d738..a7e87bc0 100644 --- a/amplifier_app_cli/commands/provider.py +++ b/amplifier_app_cli/commands/provider.py @@ -1223,8 +1223,32 @@ def provider_login(ctx: click.Context, provider_id: str) -> None: Examples: amplifier provider login openai-chatgpt """ - module_id = _normalize_module_id(provider_id) - display = _display_name(module_id) + # Resolve an exact saved instance before treating the argument as a module + # name. Login can replace credentials, so a module-only match must never + # choose among several accounts by priority or list order. + settings = _get_settings() + providers = settings.get_provider_overrides() + matches = [ + entry + for entry in providers + if entry.get("id") + and _normalize_id(entry["id"]) == _normalize_id(provider_id) + ] + if not matches: + requested_module = _normalize_module_id(provider_id) + matches = [ + entry for entry in providers + if _normalize_module_id(entry.get("module", "")) == requested_module + ] + if len(matches) > 1: + raise click.ClickException( + "More than one configured provider matches this module. " + "Choose an exact instance ID from `amplifier provider list`." + ) + entry = matches[0] if matches else None + module_id = _normalize_module_id(entry["module"] if entry else provider_id) + display = entry.get("id") if entry else None + display = display or _display_name(module_id) if not is_provider_module_installed(module_id): console.print(f"[red]Provider '{display}' is not installed.[/red]") @@ -1255,8 +1279,6 @@ def provider_login(ctx: click.Context, provider_id: str) -> None: # exists, so login uses the same connection values the wizard/runtime # would -- empty dict (not an error) when the provider has never been # configured yet. - settings = _get_settings() - entry = _find_provider_entry(settings.get_provider_overrides(), provider_id) stored_config = entry.get("config", {}) if entry else {} provider_instance = _try_instantiate_provider(provider_class, stored_config) diff --git a/amplifier_app_cli/provider_loader.py b/amplifier_app_cli/provider_loader.py index 71217f26..565b4ccf 100644 --- a/amplifier_app_cli/provider_loader.py +++ b/amplifier_app_cli/provider_loader.py @@ -7,12 +7,13 @@ import asyncio import importlib import importlib.metadata +import inspect import logging import os import sys +from copy import deepcopy from pathlib import Path -from typing import TYPE_CHECKING -from typing import Any +from typing import TYPE_CHECKING, Any from .provider_diagnostics import invoke_list_models @@ -70,9 +71,7 @@ def _load_provider_module_from_source_path(provider_id: str, source_path: Path) return module -def _load_provider_module( - provider_id: str, *, source_path: Path | None = None -) -> Any: +def _load_provider_module(provider_id: str, *, source_path: Path | None = None) -> Any: """Load a provider module. Tries entry points first, then direct import. @@ -295,7 +294,7 @@ def _try_instantiate_provider( provider_class: type, collected_config: dict[str, Any] | None = None, ) -> Any | None: - """Try to instantiate a provider class with various constructor signatures. + """Instantiate a provider using its declared constructor signature. Different providers have different constructor requirements: - Standard: (api_key, config) - Anthropic, OpenAI @@ -308,7 +307,8 @@ def _try_instantiate_provider( collected_config: Optional config values collected from user (base_url, host, etc.) Returns: - Provider instance or None if all attempts fail + Provider instance, or None when its signature cannot accept the config. + Provider constructor validation/runtime errors propagate unchanged. """ collected_config = collected_config or {} @@ -324,49 +324,47 @@ def _try_instantiate_provider( host = _resolve_env_placeholder(raw_host) or "http://localhost:11434" api_key = _resolve_env_placeholder(raw_api_key) or "" - # Common exceptions to catch during instantiation attempts: - # - TypeError: wrong argument signature - # - ValueError: invalid argument values - # - RuntimeError: some providers raise this for missing dependencies (e.g., old azure-openai) - instantiation_errors = (TypeError, ValueError, RuntimeError) - - # Approach 1: Standard (api_key, config) - Anthropic, OpenAI - try: - return provider_class(api_key=api_key, config={}) - except instantiation_errors: - pass - - # Approach 2: Azure-style (keyword-only base_url with api_key) + # Bind before construction: a TypeError *inside* a valid constructor is + # a provider failure, not permission to retry with different/default + # settings. In particular, an invalid account path or auth mode must + # never fall through to a no-argument constructor for another account. try: - return provider_class(base_url=base_url, api_key=api_key, config={}) - except instantiation_errors: - pass - - # Approach 3: VLLM-style (base_url without api_key) - try: - return provider_class(base_url=base_url, config={}) - except instantiation_errors: - pass + signature = inspect.signature(provider_class) + except (TypeError, ValueError): + logger.debug("Provider constructor signature is unavailable") + return None - # Approach 4: Ollama-style (host, config) - try: - return provider_class(host=host, config={}) - except instantiation_errors: - pass + parameters = signature.parameters + accepts_kwargs = any( + parameter.kind == inspect.Parameter.VAR_KEYWORD + for parameter in parameters.values() + ) + kwargs: dict[str, Any] = {} + if "config" in parameters or accepts_kwargs: + # Some constructors normalize nested config in place. Discovery and + # login must not mutate the caller's saved configuration or the + # values the wizard later persists. + kwargs["config"] = deepcopy(collected_config) + elif collected_config: + # Metadata-only, no-argument facades remain usable for get_info(), + # but cannot act as a configured provider for login/model discovery. + return None - # Approach 5: Just config - try: - return provider_class(config={}) - except instantiation_errors: - pass + # Include optional connection arguments as well as required ones. Trying + # api_key+config first used to skip an optional host/base_url, silently + # selecting a local/default server even after the user chose another. + for name, value in (("api_key", api_key), ("base_url", base_url), ("host", host)): + if name in parameters: + kwargs[name] = value - # Approach 6: No args try: - return provider_class() - except instantiation_errors: - pass - - return None + signature.bind(**kwargs) + except TypeError: + # A required coordinator or unknown constructor argument cannot be + # supplied by lightweight discovery. Let the caller report that the + # provider cannot be instantiated without a mounted session. + return None + return provider_class(**kwargs) def get_provider_info( diff --git a/tests/test_provider_loader_configuration.py b/tests/test_provider_loader_configuration.py new file mode 100644 index 00000000..b74e1e79 --- /dev/null +++ b/tests/test_provider_loader_configuration.py @@ -0,0 +1,334 @@ +"""Configured discovery must use the same settings as login and runtime.""" + +from copy import deepcopy +from types import SimpleNamespace +from typing import ClassVar +from unittest.mock import MagicMock, patch + +import pytest +from click.testing import CliRunner + +from amplifier_app_cli import provider_config_utils as wizard +from amplifier_app_cli import provider_loader as loader +from amplifier_app_cli.lib.settings import AppSettings, SettingsPaths + + +@pytest.fixture(autouse=True) +def isolated_settings_home(tmp_path, monkeypatch): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.setenv("AMPLIFIER_HOME", str(tmp_path / ".amplifier")) + + +class StandardProvider: + def __init__(self, api_key=None, config=None): + self.api_key, self.config = api_key, config + + +class EndpointProvider: + def __init__(self, *, base_url=None, api_key=None, config=None): + self.base_url, self.api_key, self.config = base_url, api_key, config + + +class HostProvider: + def __init__(self, host=None, config=None, api_key=None): + self.host, self.api_key, self.config = host, api_key, config + + +class ConfigProvider: + def __init__(self, config=None): + self.config = config + + +@pytest.mark.parametrize( + "provider_class", [StandardProvider, EndpointProvider, HostProvider, ConfigProvider] +) +def test_config_reaches_each_constructor_and_is_independent( + provider_class, monkeypatch +): + monkeypatch.setenv("FIXTURE_KEY", "fixture-key") + config = { + "auth_mode": "alternate", + "api_key": "${FIXTURE_KEY}", + "base_url": "https://endpoint.invalid/v1", + "host": "https://host.invalid", + "token_file_path": "/fixture/account.json", + "options": {"values": [1]}, + } + original = deepcopy(config) + provider = loader._try_instantiate_provider(provider_class, config) + assert provider.config == config + if hasattr(provider, "api_key"): + assert provider.api_key == "fixture-key" + if hasattr(provider, "base_url"): + assert provider.base_url == config["base_url"] + if hasattr(provider, "host"): + assert provider.host == config["host"] + provider.config["options"]["values"].append(2) + assert config == original + + +@pytest.mark.parametrize("error_type", [ValueError, RuntimeError, TypeError]) +def test_constructor_validation_never_retries_using_defaults(error_type): + attempts = [] + + class RejectingProvider: + def __init__(self, config=None): + attempts.append(config) + if config: + raise error_type("Invalid configured account") + + with pytest.raises(error_type, match="Invalid configured account"): + loader._try_instantiate_provider(RejectingProvider, {"auth_mode": "invalid"}) + assert attempts == [{"auth_mode": "invalid"}] + + +def test_no_arg_metadata_discovery_remains_supported_without_discarding_config(): + class MetadataOnlyProvider: + pass + + assert isinstance( + loader._try_instantiate_provider(MetadataOnlyProvider), MetadataOnlyProvider + ) + assert ( + loader._try_instantiate_provider( + MetadataOnlyProvider, {"auth_mode": "alternate"} + ) + is None + ) + + +def test_unsupported_required_coordinator_does_not_run_constructor(): + class MountedOnlyProvider: + def __init__(self, config, coordinator): + pytest.fail("Lightweight discovery cannot invent a coordinator") + + assert ( + loader._try_instantiate_provider(MountedOnlyProvider, {"account": "test"}) + is None + ) + + +def test_kwargs_provider_receives_config(): + class FlexibleProvider: + def __init__(self, **kwargs): + self.kwargs = kwargs + + assert loader._try_instantiate_provider( + FlexibleProvider, {"account": "test"} + ).kwargs == {"config": {"account": "test"}} + + +def test_catalog_uses_the_supplied_config(monkeypatch): + class CatalogProvider(ConfigProvider): + async def list_models(self): + return [SimpleNamespace(id=self.config["catalog"])] + + monkeypatch.setattr(loader, "load_provider_class", lambda _: CatalogProvider) + models = loader.get_provider_models( + "fixture", collected_config={"catalog": "selected-catalog"} + ) + assert [m.id for m in models] == ["selected-catalog"] + + +def test_azure_endpoint_alias_is_resolved(monkeypatch): + monkeypatch.setenv("FIXTURE_ENDPOINT", "https://azure.invalid") + provider = loader._try_instantiate_provider( + EndpointProvider, {"azure_endpoint": "${FIXTURE_ENDPOINT}"} + ) + assert provider.base_url == "https://azure.invalid" + + +class OAuthProvider(ConfigProvider): + events: ClassVar[list] = [] + + def get_info(self): + return SimpleNamespace( + display_name="Fixture OAuth", + capabilities=["auth:browser"], + config_fields=[ + { + "id": "auth_mode", + "display_name": "Connection", + "field_type": "choice", + "prompt": "Select connection", + "choices": ["default", "alternate"], + "default": "default", + }, + { + "id": "token_file_path", + "display_name": "Account path", + "field_type": "text", + "prompt": "Account path", + "required": False, + }, + ], + ) + + def auth_status(self): + self.events.append(("status", deepcopy(self.config))) + return "unauthenticated" + + async def login(self, print_fn=None): + self.events.append(("login", deepcopy(self.config))) + return True + + async def list_models(self): + self.events.append(("models", deepcopy(self.config))) + return [SimpleNamespace(id="chosen-model", display_name="Chosen model")] + + +def test_wizard_collects_then_reuses_configuration_for_login_and_models(monkeypatch): + OAuthProvider.events = [] + monkeypatch.setattr(loader, "load_provider_class", lambda _: OAuthProvider) + monkeypatch.setattr(wizard, "load_provider_class", lambda _: OAuthProvider) + choices = iter(["2", "/fixture/alternate.json", "1"]) + monkeypatch.setattr(wizard.Prompt, "ask", lambda *args, **kwargs: next(choices)) + monkeypatch.setattr(wizard.Confirm, "ask", lambda *args, **kwargs: True) + result = wizard.configure_provider("fixture", MagicMock()) + expected = {"auth_mode": "alternate", "token_file_path": "/fixture/alternate.json"} + assert OAuthProvider.events == [ + ("status", expected), + ("login", expected), + ("models", expected), + ] + assert result == {**expected, "default_model": "chosen-model"} + + +def test_login_command_passes_saved_configuration_through_real_loader( + tmp_path, monkeypatch +): + from amplifier_app_cli.commands.provider import provider + + settings = AppSettings( + paths=SettingsPaths( + global_settings=tmp_path / "global.yaml", + project_settings=tmp_path / "project.yaml", + local_settings=tmp_path / "local.yaml", + ) + ) + config = { + "auth_mode": "alternate", + "token_file_path": str(tmp_path / "account.json"), + } + monkeypatch.setattr(loader, "load_provider_class", lambda _: OAuthProvider) + monkeypatch.setattr(wizard, "load_provider_class", lambda _: OAuthProvider) + settings.set_provider_override( + {"module": "provider-fixture", "config": config}, scope="global" + ) + OAuthProvider.events = [] + with ( + patch( + "amplifier_app_cli.commands.provider._get_settings", return_value=settings + ), + patch( + "amplifier_app_cli.commands.provider.is_provider_module_installed", + return_value=True, + ), + patch( + "amplifier_app_cli.commands.provider.load_provider_class", + return_value=OAuthProvider, + ), + ): + result = CliRunner().invoke(provider, ["login", "fixture"]) + assert result.exit_code == 0, result.output + assert ("login", config) in OAuthProvider.events + + +@pytest.mark.parametrize("name", ["first-account", "second-account"]) +def test_login_uses_named_instance_even_when_module_has_multiple_accounts( + tmp_path, monkeypatch, name +): + from amplifier_app_cli.commands.provider import provider + + settings = AppSettings( + paths=SettingsPaths( + global_settings=tmp_path / "global.yaml", + project_settings=tmp_path / "project.yaml", + local_settings=tmp_path / "local.yaml", + ) + ) + monkeypatch.setattr(loader, "load_provider_class", lambda _: OAuthProvider) + settings._write_scope( + "global", + { + "config": { + "providers": [ + { + "id": account, + "module": "provider-fixture", + "config": {"account": account}, + } + for account in ("first-account", "second-account") + ] + } + }, + ) + OAuthProvider.events = [] + with ( + patch( + "amplifier_app_cli.commands.provider._get_settings", return_value=settings + ), + patch( + "amplifier_app_cli.commands.provider.is_provider_module_installed", + return_value=True, + ) as installed, + patch( + "amplifier_app_cli.commands.provider.load_provider_class", + return_value=OAuthProvider, + ) as load, + ): + result = CliRunner().invoke(provider, ["login", name]) + assert result.exit_code == 0, result.output + installed.assert_called_once_with("provider-fixture") + load.assert_called_once_with("provider-fixture") + assert ("login", {"account": name}) in OAuthProvider.events + assert all(config == {"account": name} for _, config in OAuthProvider.events) + + +@pytest.mark.parametrize("name", ["fixture", "provider-fixture"]) +def test_login_rejects_ambiguous_module_without_loading_or_logging_in( + tmp_path, monkeypatch, name +): + from amplifier_app_cli.commands.provider import provider + + settings = AppSettings( + paths=SettingsPaths( + global_settings=tmp_path / "global.yaml", + project_settings=tmp_path / "project.yaml", + local_settings=tmp_path / "local.yaml", + ) + ) + monkeypatch.setattr(loader, "load_provider_class", lambda _: OAuthProvider) + settings._write_scope( + "global", + { + "config": { + "providers": [ + { + "id": account, + "module": "provider-fixture", + "config": {"account": account, "priority": priority}, + } + for account, priority in ( + ("first-account", 1), + ("second-account", 99), + ) + ] + } + }, + ) + with ( + patch( + "amplifier_app_cli.commands.provider._get_settings", return_value=settings + ), + patch( + "amplifier_app_cli.commands.provider.is_provider_module_installed" + ) as installed, + patch("amplifier_app_cli.commands.provider.load_provider_class") as load, + ): + result = CliRunner().invoke(provider, ["login", name]) + assert result.exit_code == 1 + assert "Choose an exact instance ID" in result.output + installed.assert_not_called() + load.assert_not_called()