Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,12 @@ amplifier session delete <id> # Delete session
amplifier session cleanup [--days N] # Clean up old sessions
```

When a saved instance uses its provider type as its ID (for example,
`id: anthropic` with `module: provider-anthropic`), `run --provider anthropic`
applies that saved choice to a single unnamed mount of the same module instead
of adding a duplicate. Existing mounted IDs still take precedence. If several
unnamed mounts match, give the mounts unique IDs and select one explicitly.

### Conversational Single-Shot Workflows

**Build context across multiple commands:**
Expand Down
50 changes: 46 additions & 4 deletions amplifier_app_cli/commands/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from ..console import console
from ..session_handoff import takeover_options
from ..effective_config import get_effective_config_summary
from ..lib.merge_utils import merge_module_items
from ..lib.settings import AppSettings
from ..paths import create_config_manager
from ..runtime.config import resolve_config
Expand Down Expand Up @@ -467,10 +468,9 @@ def restore_diagnostic_streams() -> None:
)
providers_list = config_data.get("providers", [])

# Find the target provider — two-pass search:
# Pass 1: exact match on instance id/mount name.
# _map_id_to_instance_id copies id → instance_id without stripping id,
# so both fields co-exist on resolved entries; either leg can match.
# Find the target provider in priority order. First use an exact
# mounted instance id/mount name. _map_id_to_instance_id copies
# id → instance_id without stripping id, so either leg can match.
target_idx = None
for i, entry in enumerate(providers_list):
if isinstance(entry, dict) and (
Expand All @@ -479,6 +479,48 @@ def restore_diagnostic_streams() -> None:
target_idx = i
break

if target_idx is None:
# A named saved provider is deliberately omitted from the
# bundle mount plan unless selected. When it is selected,
# apply it to the one unnamed mount for the same module
# instead of adding a second provider.
selected_saved_provider = next(
(
entry
for entry in app_settings.get_provider_overrides()
if isinstance(entry, dict)
and entry.get("id") == provider
and entry.get("module") == provider_module
),
None,
)
if selected_saved_provider is not None:
unnamed_matches = [
i
for i, entry in enumerate(providers_list)
if isinstance(entry, dict)
and entry.get("module") == provider_module
and not entry.get("id")
and not entry.get("instance_id")
]
if len(unnamed_matches) > 1:
console.print(
f"[red]Error:[/red] Provider '{provider}' matches a saved "
f"configuration but {len(unnamed_matches)} unnamed "
f"'{provider_module}' mounts are ambiguous. "
"Select a mounted provider by id."
)
sys.exit(1)
if len(unnamed_matches) == 1:
target_idx = unnamed_matches[0]
providers_list = list(providers_list)
selected_provider = merge_module_items(
providers_list[target_idx], selected_saved_provider
)
selected_provider["id"] = provider
selected_provider["instance_id"] = provider
providers_list[target_idx] = selected_provider

if target_idx is None:
# Pass 2: fallback — module-type match (original behavior).
# Preserves single-instance usage: -p anthropic → provider-anthropic.
Expand Down
150 changes: 135 additions & 15 deletions tests/test_provider_commands.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from pathlib import Path
from unittest.mock import MagicMock, patch

import click
from click.testing import CliRunner

from amplifier_app_cli.lib.settings import AppSettings, SettingsPaths
Expand Down Expand Up @@ -1635,21 +1636,21 @@ def test_first_run_references_provider_add(self):
# ============================================================


def _make_run_cli_p_flag(captured: list) -> "click.Group":
def _make_run_cli_p_flag(captured: list, captured_mount_plan: list) -> "click.Group":
"""Create a minimal Click CLI with the run command registered.

``captured`` is cleared and replaced with the providers list that
``execute_single`` receives, i.e. after the provider selection logic has
run and the selected provider has been promoted to priority 0.
"""
import click
from amplifier_app_cli.commands.run import register_run_command
from unittest.mock import AsyncMock

cli = click.Group("test")

async def _execute_single(prompt, config_data, *args, **kwargs):
captured[:] = list(config_data.get("providers", []))
captured_mount_plan[:] = list(kwargs["prepared_bundle"].mount_plan["providers"])

register_run_command(
cli,
Expand All @@ -1662,23 +1663,30 @@ async def _execute_single(prompt, config_data, *args, **kwargs):
return cli


def _invoke_run_p_flag(providers_list: list, p_flag: str):
def _invoke_run_p_flag(
providers_list: list,
p_flag: str | None,
provider_overrides: list | None = None,
):
"""Run ``amplifier run -p <p_flag> --output-format json hello`` via CliRunner.

Mocks ``resolve_config`` to inject *providers_list* and suppresses the
update-check coroutine. Returns ``(CliResult, captured_providers)`` where
*captured_providers* contains the providers as passed to ``execute_single``
— the selected provider will have ``config["priority"] == 0``.
*captured_providers* and *captured_mount_plan* respectively contain the
providers passed to ``execute_single`` and installed on its prepared bundle.
"""
from click.testing import CliRunner
from unittest.mock import AsyncMock, MagicMock, patch

captured: list = []
cli = _make_run_cli_p_flag(captured)
captured_mount_plan: list = []
cli = _make_run_cli_p_flag(captured, captured_mount_plan)

fake_config: dict = {"providers": list(providers_list)}
fake_bundle = MagicMock()
fake_bundle.mount_plan = {"providers": list(providers_list)}
app_settings = MagicMock()
app_settings.get_provider_overrides.return_value = provider_overrides or []

with (
patch(
Expand All @@ -1689,25 +1697,34 @@ def _invoke_run_p_flag(providers_list: list, p_flag: str):
"amplifier_app_cli.commands.run.resolve_config",
return_value=(fake_config, fake_bundle),
),
patch(
"amplifier_app_cli.commands.run.AppSettings",
return_value=app_settings,
),
patch(
"amplifier_app_cli.utils.startup_checker.check_and_notify",
new_callable=AsyncMock,
),
):
runner = CliRunner()
command = ["run"]
if p_flag is not None:
command.extend(["-p", p_flag])
command.extend(["--output-format", "json", "hello"])
result = runner.invoke(
cli, ["run", "-p", p_flag, "--output-format", "json", "hello"]
cli, command
)

return result, captured
return result, captured, captured_mount_plan


class TestRunPFlag:
"""Tests for -p/--provider flag matching provider instance id/mount name.

Validates Fix 4 from UPSTREAM-FIXES.md: ``-p`` now does a two-pass search —
Pass 1 matches on ``id`` or ``instance_id``; Pass 2 falls back to module type
(``provider-{name}``) for backward compatibility.
Validates ``-p`` selection precedence: an exact mounted ``id`` or
``instance_id`` wins; an explicitly selected saved named provider may apply
to one matching unnamed mount; module-type fallback (``provider-{name}``)
preserves backward compatibility.
"""

def test_run_p_flag_matches_instance_id(self):
Expand All @@ -1724,7 +1741,7 @@ def test_run_p_flag_matches_instance_id(self):
"config": {"priority": 3},
},
]
result, captured = _invoke_run_p_flag(providers, "spark2-gemma")
result, captured, _ = _invoke_run_p_flag(providers, "spark2-gemma")

assert result.exit_code == 0, (
f"Expected success, got exit {result.exit_code}: {result.output}"
Expand Down Expand Up @@ -1754,7 +1771,7 @@ def test_run_p_flag_matches_instance_id_non_first_position(self):
"config": {"priority": 3},
},
]
result, captured = _invoke_run_p_flag(providers, "r11-gemma")
result, captured, _ = _invoke_run_p_flag(providers, "r11-gemma")

assert result.exit_code == 0, (
f"Expected success, got exit {result.exit_code}: {result.output}"
Expand All @@ -1775,7 +1792,7 @@ def test_run_p_flag_falls_back_to_module_type(self):
providers = [
{"module": "provider-anthropic", "config": {"priority": 1}},
]
result, captured = _invoke_run_p_flag(providers, "anthropic")
result, captured, _ = _invoke_run_p_flag(providers, "anthropic")

assert result.exit_code == 0, (
f"Regression: -p anthropic no longer works via module-type fallback: "
Expand Down Expand Up @@ -1814,7 +1831,7 @@ def test_run_p_flag_end_to_end_via_resolve_bundle_config(self):
assert resolved[0].get("id") == "r11-gemma"
assert resolved[0].get("instance_id") == "r11-gemma"

result, captured = _invoke_run_p_flag(resolved, "r11-gemma")
result, captured, _ = _invoke_run_p_flag(resolved, "r11-gemma")

assert result.exit_code == 0, (
f"Expected success, got exit {result.exit_code}: {result.output}"
Expand All @@ -1830,6 +1847,109 @@ def test_run_p_flag_end_to_end_via_resolve_bundle_config(self):
f"got {spark2['config']['priority']}"
)

def test_run_p_flag_applies_selected_named_provider_to_unnamed_mount(self):
"""A selected named provider replaces its matching unnamed mount."""
root_provider = {
"module": "provider-anthropic",
"config": {
"default_model": "claude-sonnet-5-5",
"priority": 5,
"use_streaming": True,
},
}
saved_provider = {
"id": "anthropic",
"module": "provider-anthropic",
"source": "git+https://example.invalid/anthropic@selected",
"config": {
"default_model": "claude-sonnet-5",
"use_streaming": False,
},
}

result, captured, mount_plan = _invoke_run_p_flag(
[root_provider], "anthropic", [saved_provider]
)

assert result.exit_code == 0, result.output
assert len(captured) == len(mount_plan) == 1
assert captured == mount_plan
selected = captured[0]
assert selected["id"] == selected["instance_id"] == "anthropic"
assert selected["source"] == saved_provider["source"]
assert selected["config"] == {
"default_model": "claude-sonnet-5",
"priority": 0,
"use_streaming": False,
}

def test_run_p_flag_rejects_ambiguous_unnamed_mounts_for_selected_provider(self):
"""Selected saved choices never silently choose between unnamed mounts."""
providers = [
{"module": "provider-anthropic", "config": {"priority": 1}},
{"module": "provider-anthropic", "config": {"priority": 2}},
]
saved_provider = {
"id": "anthropic",
"module": "provider-anthropic",
"config": {"default_model": "claude-sonnet-5"},
}

result, captured, mount_plan = _invoke_run_p_flag(
providers, "anthropic", [saved_provider]
)

assert result.exit_code == 1
assert "2 unnamed" in result.output
assert "'provider-anthropic' mounts are ambiguous" in result.output
assert captured == mount_plan == []

def test_unselected_named_provider_does_not_leak_into_unnamed_mount(self):
"""Saved named choices remain omitted until explicitly selected."""
root_provider = {
"module": "provider-anthropic",
"config": {"default_model": "claude-sonnet-5-5", "priority": 5},
}
saved_provider = {
"id": "anthropic",
"module": "provider-anthropic",
"config": {"default_model": "claude-sonnet-5"},
}

result, captured, mount_plan = _invoke_run_p_flag(
[root_provider], None, [saved_provider]
)

assert result.exit_code == 0, result.output
assert captured == mount_plan == [root_provider]

def test_run_p_flag_exact_mounted_id_wins_over_saved_provider(self):
"""An already-mounted matching id remains authoritative."""
mounted_provider = {
"id": "anthropic",
"instance_id": "anthropic",
"module": "provider-anthropic",
"config": {"default_model": "mounted-model", "priority": 5},
}
saved_provider = {
"id": "anthropic",
"module": "provider-anthropic",
"source": "git+https://example.invalid/anthropic@saved",
"config": {"default_model": "saved-model"},
}

result, captured, mount_plan = _invoke_run_p_flag(
[mounted_provider], "anthropic", [saved_provider]
)

assert result.exit_code == 0, result.output
assert captured == mount_plan
assert captured[0]["config"] == {
"default_model": "mounted-model",
"priority": 0,
}
assert "source" not in captured[0]


class TestFindProviderEntryPrioritySelection:
"""BUG 3 regression: ``_find_provider_entry()`` is the third function with
Expand Down
Loading