diff --git a/README.md b/README.md index 4cb6367d..66138727 100644 --- a/README.md +++ b/README.md @@ -129,6 +129,10 @@ responsibility to maintain by hand in `settings.yaml`; config tooling across reconfigures but never prompts for it, displays it in the wizard, or validates its contents. +Optional settings-only `auto_continue` overrides for truncated responses are +also preserved verbatim across reconfiguration, including `false`. If absent, +the CLI leaves the key unset so the provider's runtime default applies. + ### Session Commands ```bash diff --git a/amplifier_app_cli/provider_config_utils.py b/amplifier_app_cli/provider_config_utils.py index c07bfa32..19652222 100644 --- a/amplifier_app_cli/provider_config_utils.py +++ b/amplifier_app_cli/provider_config_utils.py @@ -30,17 +30,19 @@ console = Console() logger = logging.getLogger(__name__) -# Reserved, user-owned provider-config keys. Provider modules are adopting -# ``extra_request_params`` as an owner-beware dict merged verbatim into every -# API request, for parameters the module itself doesn't wrap -- users -# maintain it by hand in settings.yaml. It is deliberately never declared as -# a ConfigField (never a wizard prompt, never displayed, never validated), -# so every config-rewrite path that rebuilds a provider's config purely from -# wizard-collected fields must round-trip it verbatim via -# ``_preserve_reserved_keys()`` below rather than silently dropping it. Keep -# this tuple as the single, deliberately narrow allow-list for that -# passthrough convention -- do not widen it to cover arbitrary unknown keys. -RESERVED_PROVIDER_CONFIG_KEYS: tuple[str, ...] = ("extra_request_params",) +# Reserved, user-owned settings-only provider-config keys. Users maintain +# ``extra_request_params`` (an owner-beware dict of raw API request options) +# and ``auto_continue`` (an optional truncation-continuation override) by hand +# in settings.yaml, outside the wizard's ConfigFields. Every config-rewrite +# path that rebuilds a provider's config purely from wizard-collected fields +# must round-trip these keys verbatim via ``_preserve_reserved_keys()`` below, +# rather than silently dropping them. Keep this tuple as the single, +# deliberately narrow allow-list -- do not retain arbitrary unknown keys or +# inject defaults for absent keys. +RESERVED_PROVIDER_CONFIG_KEYS: tuple[str, ...] = ( + "extra_request_params", + "auto_continue", +) def _preserve_reserved_keys( @@ -51,9 +53,10 @@ def _preserve_reserved_keys( Called wherever a config-rewrite path replaces an EXISTING provider instance's config with a freshly wizard-collected one. Only keys in - ``RESERVED_PROVIDER_CONFIG_KEYS`` (today, just ``extra_request_params``) - are preserved -- any other non-schema key in ``old_config`` is dropped, - exactly as before this function existed. If ``new_config`` already + ``RESERVED_PROVIDER_CONFIG_KEYS`` (settings-only keys such as + ``extra_request_params`` and ``auto_continue``) are preserved -- any + other non-schema key in ``old_config`` is dropped, exactly as before + this function existed. If ``new_config`` already carries the key (e.g. a module surfaces it deliberately in the future), the old value is not used -- the freshly collected value wins. diff --git a/tests/test_provider_config_reserved_keys.py b/tests/test_provider_config_reserved_keys.py index d264be7e..e62485cc 100644 --- a/tests/test_provider_config_reserved_keys.py +++ b/tests/test_provider_config_reserved_keys.py @@ -2,21 +2,21 @@ Provider modules are adopting an ``extra_request_params`` config key -- an owner-beware dict merged verbatim into every API request, for parameters the -module itself doesn't wrap. Users maintain this key by hand in -settings.yaml; it is never declared as a ``ConfigField`` by any provider -module and the wizard never prompts for it. +module itself doesn't wrap. Users maintain this key and optional +``auto_continue`` truncation-continuation overrides by hand in settings.yaml, +outside the wizard's ``ConfigField`` schema. Bug: every edit/reconfigure flow rebuilds a provider's config purely from the module's declared ``ConfigField`` schema answers (``configure_provider()`` in ``provider_config_utils.py``), then callers assign the rebuilt dict wholesale over the old one. Any key not in the schema -- including -``extra_request_params`` -- is silently DROPPED on edit, defeating the whole -point of a hand-maintained passthrough bag. +``extra_request_params`` and ``auto_continue`` -- is silently DROPPED on edit, +defeating settings-only overrides. Fix: ``_preserve_reserved_keys(old, new)`` in provider_config_utils.py -carries ``extra_request_params`` forward verbatim from the prior config into -the rebuilt one, if present, at every seam that replaces an EXISTING -provider instance's config: +carries reserved keys forward verbatim from the prior config into the rebuilt +one, if present, at every seam that replaces an EXISTING provider instance's +config: * ``provider_edit()`` (commands/provider.py) * ``_manage_edit_provider()`` (commands/provider.py, interactive loop) @@ -32,11 +32,17 @@ empty config_fields schema (d) a fresh add (no prior config) never invents the key (e) unit coverage of _preserve_reserved_keys() itself + (f) native False and legacy string values retain their exact types + (g) all four CLI rewrite seams preserve a settings-only auto_continue + (h) real interactive/non-interactive wizard returns leave reserved keys + to the caller's preservation seam, without inventing defaults """ +from copy import deepcopy from pathlib import Path from unittest.mock import MagicMock, patch +import pytest from click.testing import CliRunner import amplifier_app_cli.provider_config_utils as pcu @@ -370,3 +376,319 @@ def test_only_the_declared_reserved_keys_are_preserved(self): result = pcu._preserve_reserved_keys(old, new) assert result == {"extra_request_params": {"a": 1}} assert "some_stale_key" not in result + + @pytest.mark.parametrize("value", [False, "false", True]) + def test_preserves_auto_continue_verbatim_without_mutating_inputs(self, value): + old = { + "default_model": "old-model", + "auto_continue": value, + "extra_request_params": {"service_tier": "fast"}, + "ghost": "obsolete", + } + new = {"default_model": "new-model"} + old_snapshot, new_snapshot = deepcopy(old), deepcopy(new) + + result = pcu._preserve_reserved_keys(old, new) + + assert result == { + "default_model": "new-model", + "auto_continue": value, + "extra_request_params": {"service_tier": "fast"}, + } + assert type(result["auto_continue"]) is type(value) + assert "ghost" not in result + assert old == old_snapshot + assert new == new_snapshot + + @pytest.mark.parametrize("old_value", [False, "false", True]) + @pytest.mark.parametrize("new_value", [False, "false", True]) + def test_explicit_new_auto_continue_wins_without_mutating_inputs( + self, old_value, new_value + ): + old = {"auto_continue": old_value, "extra_request_params": {"a": 1}} + new = {"auto_continue": new_value} + old_snapshot, new_snapshot = deepcopy(old), deepcopy(new) + + result = pcu._preserve_reserved_keys(old, new) + + assert result == {"auto_continue": new_value, "extra_request_params": {"a": 1}} + assert type(result["auto_continue"]) is type(new_value) + assert old == old_snapshot + assert new == new_snapshot + + @pytest.mark.parametrize("old", [None, {}, {"default_model": "old-model"}]) + def test_missing_auto_continue_is_not_defaulted(self, old): + new = {"default_model": "new-model"} + result = pcu._preserve_reserved_keys(old, new) + assert result == new + assert "auto_continue" not in result + + +@pytest.fixture +def settings_only_provider_info(): + """Four ordinary optional fields, no credentials or settings-only fields.""" + return { + "display_name": "Test Provider", + "capabilities": [], + "config_fields": [ + { + "id": "base_url", + "display_name": "API Base URL", + "prompt": "API base URL", + "field_type": "text", + "required": False, + }, + { + "id": "max_tokens", + "display_name": "Max Output Tokens", + "prompt": "Max output tokens", + "field_type": "text", + "required": False, + "requires_model": True, + }, + { + "id": "temperature", + "display_name": "Temperature", + "prompt": "Sampling temperature", + "field_type": "text", + "required": False, + "requires_model": True, + }, + { + "id": "streaming", + "display_name": "Streaming", + "prompt": "Stream responses?", + "field_type": "boolean", + "required": False, + "requires_model": True, + }, + ], + } + + +@pytest.fixture +def isolated_provider_commands(tmp_path, monkeypatch, settings_only_provider_info): + """Use real persistence/preservation, isolating discovery and credentials.""" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.setenv("AMPLIFIER_HOME", str(tmp_path / ".amplifier")) + monkeypatch.setattr(Path, "home", lambda: tmp_path) + monkeypatch.chdir(tmp_path) + settings = _make_settings(tmp_path) + + with ( + patch( + "amplifier_app_cli.commands.provider._get_settings", + return_value=settings, + ), + patch("amplifier_app_cli.commands.provider._ensure_providers_ready"), + patch("amplifier_app_cli.commands.provider.KeyManager"), + patch("amplifier_app_cli.commands.provider.ProviderManager") as manager, + patch.object(pcu, "get_provider_info", return_value=settings_only_provider_info), + ): + manager.return_value.list_providers.return_value = [ + ("provider-test-provider", "Test Provider", "Test metadata") + ] + yield settings + + +class TestAutoContinueRewriteSeams: + @pytest.mark.parametrize( + ("args", "user_input", "replace_during_wizard"), + [ + pytest.param(["edit", "test-provider"], "", False, id="edit"), + pytest.param(["manage"], "e1\nd\n", False, id="manage-edit"), + pytest.param(["add", "test-provider"], "", True, id="add-replace"), + pytest.param(["manage"], "a\n1\nd\n", True, id="manage-add-replace"), + ], + ) + def test_cli_rewrites_preserve_native_false( + self, isolated_provider_commands, args, user_input, replace_during_wizard + ): + from amplifier_app_cli.commands.provider import provider + + settings = isolated_provider_commands + old = { + "default_model": "old-model", + "auto_continue": False, + "extra_request_params": {"service_tier": "fast"}, + "ghost": "obsolete", + } + old_snapshot = deepcopy(old) + if not replace_during_wizard: + _seed_provider(settings, "provider-test-provider", old) + wizard_answer = {"default_model": "new-model"} + + def rebuild_config(*args, **kwargs): + if replace_during_wizard: + # Exercise the real add-without-id replacement seam: a + # concurrent same-module entry lands after the pre-lock read. + # Normal same-module adds instead create a separate instance. + _seed_provider(settings, "provider-test-provider", old) + else: + assert kwargs["existing_config"]["auto_continue"] is False + return wizard_answer + + with patch( + "amplifier_app_cli.commands.provider.configure_provider", + side_effect=rebuild_config, + ) as configure: + result = CliRunner().invoke(provider, args, input=user_input) + + assert result.exit_code == 0, result.output + configure.assert_called_once() + entries = settings.get_scope_provider_overrides("global") + assert len(entries) == 1 + assert entries[0]["module"] == "provider-test-provider" + assert "id" not in entries[0] + config = entries[0]["config"] + assert config["default_model"] == "new-model" + assert config["auto_continue"] is False + assert config["extra_request_params"] == {"service_tier": "fast"} + assert "ghost" not in config + assert wizard_answer == {"default_model": "new-model"} + assert old == old_snapshot + + @pytest.mark.parametrize( + ("args", "user_input"), + [ + pytest.param(["add", "test-provider"], "", id="add"), + pytest.param(["manage"], "a\n1\nd\n", id="manage-add"), + ], + ) + def test_cli_fresh_add_leaves_auto_continue_absent( + self, isolated_provider_commands, args, user_input + ): + from amplifier_app_cli.commands.provider import provider + + settings = isolated_provider_commands + with patch( + "amplifier_app_cli.commands.provider.configure_provider", + return_value={"default_model": "new-model"}, + ) as configure: + result = CliRunner().invoke(provider, args, input=user_input) + + assert result.exit_code == 0, result.output + configure.assert_called_once() + entries = settings.get_scope_provider_overrides("global") + assert len(entries) == 1 + assert entries[0]["config"]["default_model"] == "new-model" + assert "auto_continue" not in entries[0]["config"] + assert "extra_request_params" not in entries[0]["config"] + + +class TestAutoContinueWizardBoundary: + @pytest.mark.parametrize("reconfigure", [False, True], ids=["fresh", "edit"]) + def test_interactive_wizard_does_not_prompt_or_return_auto_continue( + self, isolated_provider_commands, reconfigure + ): + old = ( + { + "default_model": "old-model", + "auto_continue": False, + "extra_request_params": {"service_tier": "fast"}, + "ghost": "obsolete", + } + if reconfigure + else None + ) + old_snapshot = deepcopy(old) + key_manager = MagicMock() + model_configs = [] + + def answer_text(prompt, **kwargs): + return "new-model" if prompt == "Model name" else kwargs.get("default", "") + + def list_models(provider_id, collected_config): + # Capture before configure_provider adds the selected model to + # this mutable dict; mock call_args retain the same reference. + model_configs.append((provider_id, deepcopy(collected_config))) + return [] + + with ( + patch.object(pcu, "get_provider_models", side_effect=list_models) as models, + patch.object(pcu, "console", MagicMock()), + patch.object(pcu.Prompt, "ask", side_effect=answer_text) as prompt, + patch.object( + pcu.Confirm, "ask", side_effect=lambda *a, **kw: kw.get("default") + ) as confirm, + ): + collected = pcu.configure_provider( + "provider-test-provider", + key_manager, + existing_config=old, + settings=isolated_provider_commands, + ) + + assert collected == {"default_model": "new-model"} + assert [call.args[0] for call in prompt.call_args_list] == [ + "API base URL", + "Model name", + "Max output tokens", + "Sampling temperature", + ] + assert confirm.call_count == 1 + assert confirm.call_args.args[0].startswith("Stream responses?") + for call in prompt.call_args_list + confirm.call_args_list: + text = " ".join(str(arg) for arg in call.args).lower() + assert "continue truncated" not in text + assert "auto" not in text + models.assert_called_once() + assert model_configs == [("test-provider", {})] + key_manager.save_key.assert_not_called() + + preserved = pcu._preserve_reserved_keys(old, collected) + if reconfigure: + assert preserved["auto_continue"] is False + assert preserved["extra_request_params"] == {"service_tier": "fast"} + else: + assert "auto_continue" not in preserved + assert "extra_request_params" not in preserved + assert "ghost" not in preserved + assert collected == {"default_model": "new-model"} + assert old == old_snapshot + + def test_non_interactive_wizard_return_preserves_override_at_caller_seam( + self, isolated_provider_commands + ): + old = { + "base_url": "https://api.example.test", + "default_model": "old-model", + "max_tokens": 128, + "temperature": "0.5", + "streaming": False, + "auto_continue": False, + "extra_request_params": {"service_tier": "fast"}, + "ghost": "obsolete", + } + old_snapshot = deepcopy(old) + with ( + patch.object(pcu, "get_provider_models") as models, + patch.object(pcu.Prompt, "ask") as prompt, + patch.object(pcu.Confirm, "ask") as confirm, + ): + collected = pcu.configure_provider( + "provider-test-provider", + MagicMock(), + existing_config=old, + non_interactive=True, + settings=isolated_provider_commands, + ) + + assert collected == { + "base_url": "https://api.example.test", + "default_model": "old-model", + "max_tokens": 128, + "temperature": "0.5", + "streaming": False, + } + prompt.assert_not_called() + confirm.assert_not_called() + models.assert_not_called() + collected_snapshot = deepcopy(collected) + preserved = pcu._preserve_reserved_keys(old, collected) + assert preserved["auto_continue"] is False + assert preserved["extra_request_params"] == {"service_tier": "fast"} + assert "ghost" not in preserved + assert collected == collected_snapshot + assert old == old_snapshot