diff --git a/amplifier_app_cli/effective_config.py b/amplifier_app_cli/effective_config.py index b9737978..5fe417e7 100644 --- a/amplifier_app_cli/effective_config.py +++ b/amplifier_app_cli/effective_config.py @@ -5,12 +5,9 @@ from __future__ import annotations -import logging from dataclasses import dataclass from typing import Any -logger = logging.getLogger(__name__) - @dataclass class EffectiveConfigSummary: @@ -128,23 +125,16 @@ def _select_provider_by_priority( def _get_provider_display_name(provider_module: str) -> str: """Get friendly display name for a provider module. + Banner rendering must not import provider implementations before the + session loader selects their configured source. + Args: provider_module: Provider module ID (e.g., "provider-azure-openai") Returns: Friendly display name (e.g., "Azure OpenAI") """ - # Try to get from provider's get_info() - try: - from .provider_loader import get_provider_info - - info = get_provider_info(provider_module) - if info and "display_name" in info: - return info["display_name"] - except Exception as e: - logger.debug(f"Could not get provider info for {provider_module}: {e}") - - # Fallback: Convert module ID to friendly name + # Convert the module ID without loading a provider for display metadata. # "provider-azure-openai" -> "Azure OpenAI" name = provider_module.replace("provider-", "") # Handle common cases diff --git a/amplifier_app_cli/provider_loader.py b/amplifier_app_cli/provider_loader.py index fffc3e71..71217f26 100644 --- a/amplifier_app_cli/provider_loader.py +++ b/amplifier_app_cli/provider_loader.py @@ -9,6 +9,8 @@ import importlib.metadata import logging import os +import sys +from pathlib import Path from typing import TYPE_CHECKING from typing import Any @@ -36,7 +38,41 @@ def _get_provider_module_name(provider_id: str) -> str: return f"amplifier_module_provider_{provider_id.replace('-', '_')}" -def _load_provider_module(provider_id: str) -> Any: +def _load_provider_module_from_source_path(provider_id: str, source_path: Path) -> Any: + """Import a provider from its activated source and verify its origin. + + ``BundleModuleResolver`` activates a module by adding its source root to + ``sys.path``. Import through that normal mechanism, but refuse an already + cached package or an import result that belongs to another root. + """ + module_name = _get_provider_module_name(provider_id) + expected_file = (source_path / module_name / "__init__.py").resolve() + if not expected_file.is_file(): + raise ImportError( + f"Configured source for provider '{provider_id}' does not contain " + f"expected package '{module_name}'" + ) + + cached_module = sys.modules.get(module_name) + if cached_module is not None: + cached_file = getattr(cached_module, "__file__", None) + if cached_file is None or Path(cached_file).resolve() != expected_file: + raise ImportError( + f"Provider '{provider_id}' is already loaded from a different source" + ) + + module = importlib.import_module(module_name) + loaded_file = getattr(module, "__file__", None) + if loaded_file is None or Path(loaded_file).resolve() != expected_file: + raise ImportError( + f"Provider '{provider_id}' loaded from a different source than configured" + ) + return module + + +def _load_provider_module( + provider_id: str, *, source_path: Path | None = None +) -> Any: """Load a provider module. Tries entry points first, then direct import. @@ -50,6 +86,9 @@ def _load_provider_module(provider_id: str) -> Any: Raises: ImportError: If module cannot be loaded """ + if source_path is not None: + return _load_provider_module_from_source_path(provider_id, source_path) + # Normalize to full module ID module_id = ( provider_id @@ -76,7 +115,9 @@ def _load_provider_module(provider_id: str) -> Any: raise ImportError(f"Could not load provider module '{provider_id}': {e}") from e -def load_provider_class(provider_id: str) -> type | None: +def load_provider_class( + provider_id: str, *, source_path: Path | None = None +) -> type | None: """Load a provider class for configuration purposes. This is a lightweight load that doesn't require a full coordinator. @@ -85,12 +126,16 @@ def load_provider_class(provider_id: str) -> type | None: Args: provider_id: Provider ID (e.g., "provider-anthropic" or "anthropic") + source_path: Activated source root to load instead of discovery. Returns: Provider class if found, None otherwise """ try: - module = _load_provider_module(provider_id) + if source_path is None: + module = _load_provider_module(provider_id) + else: + module = _load_provider_module(provider_id, source_path=source_path) # Look for provider class in module's __all__ or by convention # Convention: {Name}Provider (e.g., AnthropicProvider) @@ -124,6 +169,8 @@ def load_provider_class(provider_id: str) -> type | None: return None except ImportError as e: + if source_path is not None: + raise logger.debug(f"Could not load provider class for '{provider_id}': {e}") return None @@ -322,17 +369,24 @@ def _try_instantiate_provider( return None -def get_provider_info(provider_id: str) -> dict[str, Any] | None: +def get_provider_info( + provider_id: str, *, source_path: Path | None = None +) -> dict[str, Any] | None: """Get provider metadata. Args: provider_id: Provider ID (e.g., "provider-anthropic" or "anthropic") + source_path: Activated source root to load instead of entry-point + discovery. An explicit source must load from that exact root. Returns: Provider info dict if available, None otherwise """ try: - provider_class = load_provider_class(provider_id) + if source_path is None: + provider_class = load_provider_class(provider_id) + else: + provider_class = load_provider_class(provider_id, source_path=source_path) if not provider_class: logger.debug( f"get_provider_info: load_provider_class returned None for '{provider_id}'" @@ -358,6 +412,8 @@ def get_provider_info(provider_id: str) -> dict[str, Any] | None: return info.model_dump() if hasattr(info, "model_dump") else vars(info) except Exception as e: + if source_path is not None: + raise logger.warning( f"get_provider_info failed for '{provider_id}': {type(e).__name__}: {e}" ) diff --git a/amplifier_app_cli/runtime/config.py b/amplifier_app_cli/runtime/config.py index 79a9e182..a752e4ad 100644 --- a/amplifier_app_cli/runtime/config.py +++ b/amplifier_app_cli/runtime/config.py @@ -371,7 +371,11 @@ def _on_progress(action: str, detail: str) -> None: # for the reuse-or-separate multi-instance flow. raw_providers = bundle_config.get("providers") if isinstance(raw_providers, list): - _validate_provider_credentials(raw_providers) + await _validate_provider_credentials( + raw_providers, + prepared_resolver=prepared.resolver, + configured_sources=combined_sources, + ) # Expand environment variables # IMPORTANT: Must expand BEFORE syncing to mount_plan, so ${ANTHROPIC_API_KEY} etc. become actual values @@ -1215,7 +1219,12 @@ def _merge_module_lists( ENV_PATTERN = re.compile(r"\$\{([^}:]+)(?::([^}]*))?}") -def _validate_provider_credentials(providers: list[Any]) -> None: +async def _validate_provider_credentials( + providers: list[Any], + *, + prepared_resolver: Any, + configured_sources: dict[str, Any] | None = None, +) -> None: """Fail loudly, before session mount, when a provider instance's configured credential placeholder resolves to nothing. @@ -1238,11 +1247,16 @@ def _validate_provider_credentials(providers: list[Any]) -> None: ``${VAR:-default}`` default is also left alone (the default already covers "unset"). - App-CLI policy only -- no core, provider-contract, or settings-schema - changes. When provider metadata can't be loaded (custom/removed - provider, import error, etc.), validation is skipped for that entry -- - consistent with how the rest of this module already treats a missing - ``get_provider_info()`` result. + Metadata is loaded only after a configured value is found to be a bare, + unset placeholder, so normal literal/set/default/empty configurations do + not eagerly import providers. An explicit provider source is resolved + through the prepared bundle's lazy resolver and its metadata must come + from that exact activated root. That fail-closed boundary prevents an + ambient installed provider from standing in for the configured source. + The effective source mapping passed to bundle preparation is also honored, + so module and override source configuration has the same boundary. + Without a source, metadata lookup retains the historical fail-soft + behavior: unavailable metadata skips validation for that entry. """ for entry in providers: if not isinstance(entry, dict): @@ -1252,7 +1266,34 @@ def _validate_provider_credentials(providers: list[Any]) -> None: if not isinstance(module_id, str) or not isinstance(config, dict): continue - info = get_provider_info(module_id) + has_unset_placeholder = any( + isinstance(value, str) + and (match := ENV_PATTERN.fullmatch(value)) is not None + and match.group(2) is None + and not os.environ.get(match.group(1)) + for value in config.values() + ) + if not has_unset_placeholder: + continue + + source_hint = (configured_sources or {}).get(module_id) or entry.get("source") + if source_hint: + label = entry.get("id") or module_id + try: + source = await prepared_resolver.async_resolve( + module_id, source_hint=source_hint + ) + info = get_provider_info(module_id, source_path=source.resolve()) + if not info: + raise RuntimeError("provider metadata is unavailable") + except Exception as exc: + raise ValueError( + f"Could not load configured provider source for '{label}' " + "while validating credentials. Check the configured module " + "source before starting the session." + ) from exc + else: + info = get_provider_info(module_id) if not info: continue diff --git a/tests/test_runtime_credential_validation.py b/tests/test_runtime_credential_validation.py index 4a69148d..f69b12e1 100644 --- a/tests/test_runtime_credential_validation.py +++ b/tests/test_runtime_credential_validation.py @@ -9,6 +9,11 @@ instance through the WRONG account's key. This must fail loudly instead. """ +from pathlib import Path +import json +import os +import subprocess +import sys from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -32,7 +37,8 @@ def _provider_info(env_var: str = "OPENAI_API_KEY", required: bool = True) -> di class TestValidateProviderCredentials: - def test_unset_separate_binding_fails_loud(self, monkeypatch): + @pytest.mark.asyncio + async def test_unset_separate_binding_fails_loud(self, monkeypatch): """The core hardening case: a required secret field whose ${VAR} placeholder resolves to nothing must raise BEFORE session mount, instead of silently letting expand_env_vars turn it into "" (which @@ -53,9 +59,12 @@ def test_unset_separate_binding_fails_loud(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="OPENAI_WORK_API_KEY"), ), pytest.raises(ValueError, match="OPENAI_WORK_API_KEY"): - _validate_provider_credentials(providers) + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) - def test_set_credential_does_not_raise(self, monkeypatch): + @pytest.mark.asyncio + async def test_set_credential_does_not_raise(self, monkeypatch): """A required secret field whose env var IS set must pass silently.""" monkeypatch.setenv("OPENAI_WORK_API_KEY", "sk-distinct-value") @@ -71,9 +80,12 @@ def test_set_credential_does_not_raise(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="OPENAI_WORK_API_KEY"), ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_shared_binding_with_set_var_does_not_raise(self, monkeypatch): + @pytest.mark.asyncio + async def test_shared_binding_with_set_var_does_not_raise(self, monkeypatch): """Two instances sharing the SAME credential env var (the "reuse" path) must not raise as long as that shared var is set -- this is the normal, intended shared-credential configuration.""" @@ -96,9 +108,12 @@ def test_shared_binding_with_set_var_does_not_raise(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="ANTHROPIC_API_KEY"), ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_optional_keyless_secret_field_unset_does_not_raise(self, monkeypatch): + @pytest.mark.asyncio + async def test_optional_keyless_secret_field_unset_does_not_raise(self, monkeypatch): """A secret field declared `required=False` (e.g. a local/keyless Chat Completions server) must be left alone even when unset -- the fail-loud guard is scoped to REQUIRED credential fields only.""" @@ -116,9 +131,12 @@ def test_optional_keyless_secret_field_unset_does_not_raise(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="LOCAL_SERVER_API_KEY", required=False), ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_missing_provider_metadata_skips_validation(self, monkeypatch): + @pytest.mark.asyncio + async def test_missing_provider_metadata_skips_validation(self, monkeypatch): """When provider metadata can't be loaded (custom/removed provider, import error, etc.) validation is skipped for that entry rather than blocking session start -- consistent with how the rest of @@ -137,9 +155,12 @@ def test_missing_provider_metadata_skips_validation(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=None, ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_literal_value_is_not_validated(self, monkeypatch): + @pytest.mark.asyncio + async def test_literal_value_is_not_validated(self, monkeypatch): """A literal (non-placeholder) config value is untouched by this guard -- it isn't an unresolved env var reference at all.""" providers = [ @@ -154,9 +175,12 @@ def test_literal_value_is_not_validated(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="OPENAI_WORK_API_KEY"), ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_inline_default_placeholder_is_not_validated(self, monkeypatch): + @pytest.mark.asyncio + async def test_inline_default_placeholder_is_not_validated(self, monkeypatch): """A placeholder carrying an inline default (${VAR:default}) already has its own "unset" handling via expand_env_vars -- this guard only concerns itself with bare ${VAR} placeholders.""" @@ -174,9 +198,12 @@ def test_inline_default_placeholder_is_not_validated(self, monkeypatch): "amplifier_app_cli.runtime.config.get_provider_info", return_value=_provider_info(env_var="OPENAI_WORK_API_KEY"), ): - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise - def test_multiple_providers_identifies_the_failing_instance(self, monkeypatch): + @pytest.mark.asyncio + async def test_multiple_providers_identifies_the_failing_instance(self, monkeypatch): """With several configured providers, the error must name the instance and env var actually at fault, not a generic message.""" monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-fine") @@ -204,12 +231,15 @@ def _info(module_id: str): "amplifier_app_cli.runtime.config.get_provider_info", side_effect=_info, ), pytest.raises(ValueError) as exc_info: - _validate_provider_credentials(providers) + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) assert "OPENAI_WORK_API_KEY" in str(exc_info.value) assert "openai-work" in str(exc_info.value) - def test_non_dict_or_non_module_entries_are_skipped(self): + @pytest.mark.asyncio + async def test_non_dict_or_non_module_entries_are_skipped(self): """Malformed provider entries must not crash the guard.""" providers = [ "not-a-dict", @@ -217,7 +247,152 @@ def test_non_dict_or_non_module_entries_are_skipped(self): {"module": 123, "config": {"api_key": "${SOMETHING}"}}, # bad module type {"module": "provider-x", "config": "not-a-dict"}, # bad config type ] - _validate_provider_credentials(providers) # must not raise + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) # must not raise + + @pytest.mark.asyncio + async def test_source_with_set_credential_skips_resolution_and_metadata( + self, monkeypatch + ): + """Configured sources stay lazy when no unresolved value needs checking.""" + monkeypatch.setenv("EXAMPLE_API_KEY", "set") + resolver = MagicMock() + resolver.async_resolve = AsyncMock() + providers = [ + { + "module": "provider-example", + "source": "configured-source", + "config": {"api_key": "${EXAMPLE_API_KEY}"}, + } + ] + + with patch( + "amplifier_app_cli.runtime.config.get_provider_info" + ) as get_info: + await _validate_provider_credentials( + providers, prepared_resolver=resolver + ) + + resolver.async_resolve.assert_not_awaited() + get_info.assert_not_called() + + @pytest.mark.asyncio + async def test_configured_source_mapping_wins_over_entry_source( + self, monkeypatch, tmp_path + ): + """A configured source is resolved before its required key is checked.""" + monkeypatch.delenv("EXAMPLE_WORK_API_KEY", raising=False) + source = MagicMock() + source.resolve.return_value = tmp_path + resolver = MagicMock() + resolver.async_resolve = AsyncMock(return_value=source) + providers = [ + { + "module": "provider-example", + "id": "example-work", + "source": "entry-source", + "config": {"api_key": "${EXAMPLE_WORK_API_KEY}"}, + } + ] + + with patch( + "amplifier_app_cli.runtime.config.get_provider_info", + return_value=_provider_info(env_var="EXAMPLE_WORK_API_KEY"), + ) as get_info, pytest.raises(ValueError, match="EXAMPLE_WORK_API_KEY"): + await _validate_provider_credentials( + providers, + prepared_resolver=resolver, + configured_sources={"provider-example": "mapped-source"}, + ) + + resolver.async_resolve.assert_awaited_once_with( + "provider-example", source_hint="mapped-source" + ) + get_info.assert_called_once_with("provider-example", source_path=tmp_path) + + @pytest.mark.asyncio + async def test_source_metadata_failure_does_not_fall_back_to_ambient( + self, monkeypatch, tmp_path + ): + """A source-specific metadata failure is actionable instead of fail-soft.""" + monkeypatch.delenv("EXAMPLE_WORK_API_KEY", raising=False) + source = MagicMock() + source.resolve.return_value = tmp_path + resolver = MagicMock() + resolver.async_resolve = AsyncMock(return_value=source) + providers = [ + { + "module": "provider-example", + "source": "configured-source", + "config": {"api_key": "${EXAMPLE_WORK_API_KEY}"}, + } + ] + + with patch( + "amplifier_app_cli.runtime.config.get_provider_info", + side_effect=ImportError("loaded from a different source"), + ) as get_info, pytest.raises( + ValueError, match="Could not load configured provider source" + ): + await _validate_provider_credentials( + providers, prepared_resolver=resolver + ) + + get_info.assert_called_once_with("provider-example", source_path=tmp_path) + + @pytest.mark.asyncio + async def test_no_source_preserves_fail_soft_metadata_lookup( + self, monkeypatch + ): + """Legacy entries look up metadata without a source-path argument.""" + monkeypatch.delenv("EXAMPLE_WORK_API_KEY", raising=False) + providers = [ + { + "module": "provider-example", + "config": {"api_key": "${EXAMPLE_WORK_API_KEY}"}, + } + ] + + with patch( + "amplifier_app_cli.runtime.config.get_provider_info", return_value=None + ) as get_info: + await _validate_provider_credentials( + providers, prepared_resolver=MagicMock() + ) + + get_info.assert_called_once_with("provider-example") + + @pytest.mark.asyncio + async def test_optional_unset_secret_is_allowed_after_source_metadata( + self, monkeypatch, tmp_path + ): + """Optional secret fields remain allowed after canonical source lookup.""" + monkeypatch.delenv("OPTIONAL_EXAMPLE_API_KEY", raising=False) + source = MagicMock() + source.resolve.return_value = tmp_path + resolver = MagicMock() + resolver.async_resolve = AsyncMock(return_value=source) + providers = [ + { + "module": "provider-example", + "source": "configured-source", + "config": {"api_key": "${OPTIONAL_EXAMPLE_API_KEY}"}, + } + ] + + with patch( + "amplifier_app_cli.runtime.config.get_provider_info", + return_value=_provider_info( + env_var="OPTIONAL_EXAMPLE_API_KEY", required=False + ), + ) as get_info: + await _validate_provider_credentials( + providers, prepared_resolver=resolver + ) + + resolver.async_resolve.assert_awaited_once() + get_info.assert_called_once_with("provider-example", source_path=tmp_path) # ============================================================ @@ -274,3 +449,405 @@ async def test_resolve_bundle_config_raises_for_unset_separate_credential( ),pytest.raises(ValueError, match="OPENAI_WORK_API_KEY") ): await resolve_bundle_config(bundle_name="test", app_settings=settings) + + +def _write_standard_provider(source_root: Path, *, required: bool) -> None: + """Write a source-root provider package accepted by core's real validator.""" + package = source_root / "amplifier_module_provider_example" + package.mkdir(parents=True) + package.joinpath("__init__.py").write_text( + f""" +import os +from pathlib import Path + +from amplifier_core import ConfigField, ProviderInfo + +__amplifier_module_type__ = "provider" + + +class ExampleProvider: + name = "example" + + def get_info(self): + return ProviderInfo( + id="example", + display_name="Example", + config_fields=[ + ConfigField( + id="api_key", + display_name="API key", + field_type="secret", + prompt="API key", + env_var="EXAMPLE_WORK_API_KEY", + required={required!r}, + ) + ], + ) + + async def list_models(self): + return [] + + async def complete(self, request): + raise AssertionError("No provider request allowed in this test") + + def parse_tool_calls(self, response): + return [] + + +async def mount(coordinator, config): + Path(os.environ["SOURCE_MARKER"]).write_text("mounted") + await coordinator.mount("providers", ExampleProvider(), name="example") +""".lstrip(), + encoding="utf-8", + ) + _write_session_modules(source_root) + + +def _write_session_modules(source_root: Path) -> None: + """Write local modules satisfying the real session initialization contracts.""" + (source_root / "amplifier_module_orchestrator_example").mkdir() + (source_root / "amplifier_module_orchestrator_example" / "__init__.py").write_text( + """ +__amplifier_module_type__ = "orchestrator" + +class ExampleOrchestrator: + async def execute(self, prompt, context, providers, tools, hooks): + raise AssertionError("No orchestrator execution allowed in this test") + +async def mount(coordinator, config): + await coordinator.mount("orchestrator", ExampleOrchestrator()) +""".lstrip(), + encoding="utf-8", + ) + (source_root / "amplifier_module_context_example").mkdir() + (source_root / "amplifier_module_context_example" / "__init__.py").write_text( + """ +__amplifier_module_type__ = "context" + +class ExampleContext: + async def add_message(self, message): pass + async def get_messages_for_request(self): return [] + async def get_messages(self): return [] + async def set_messages(self, messages): pass + async def clear(self): pass + +async def mount(coordinator, config): + await coordinator.mount("context", ExampleContext()) +""".lstrip(), + encoding="utf-8", + ) + + +def _write_ambient_entry_point(ambient_root: Path) -> None: + """Write a conflicting installed package and entry point that stay unused.""" + ambient_root.mkdir() + package = ambient_root / "amplifier_module_provider_example" + package.mkdir() + package.joinpath("__init__.py").write_text( + """ +import os +from pathlib import Path + +Path(os.environ["AMBIENT_MARKER"]).write_text("imported") + + +async def mount(coordinator, config): + raise AssertionError("ambient entry point must not be loaded") +""".lstrip(), + encoding="utf-8", + ) + dist_info = ambient_root / "ambient_example-0.0.dist-info" + dist_info.mkdir() + (dist_info / "METADATA").write_text( + "Metadata-Version: 2.1\nName: ambient-example\nVersion: 0.0\n", + encoding="utf-8", + ) + (dist_info / "entry_points.txt").write_text( + "[amplifier.modules]\n" + "provider-example = amplifier_module_provider_example:mount\n", + encoding="utf-8", + ) + + +@pytest.mark.parametrize("source_on_entry", [False, True]) +def test_banner_summary_does_not_import_provider_implementation( + tmp_path, source_on_entry +): + """Display must not materialize an ambient provider before source activation.""" + source_root = tmp_path / "configured-source" + ambient_root = tmp_path / "ambient" + ambient_marker = tmp_path / "ambient-marker" + _write_standard_provider(source_root, required=False) + _write_ambient_entry_point(ambient_root) + + script = r""" +import json +import sys + +from amplifier_app_cli.effective_config import get_effective_config_summary + +sys.path.insert(0, sys.argv[1]) +entry = { + "module": "provider-example", + "config": {"default_model": "fixture-model"}, +} +if sys.argv[2] == "entry": + entry["source"] = sys.argv[3] +module_name = "amplifier_module_provider_example" +assert module_name not in sys.modules +summary = get_effective_config_summary( + {"providers": [entry]}, config_source="bundle:fixture" +) +assert module_name not in sys.modules +print(json.dumps({ + "provider_name": summary.provider_name, + "banner": summary.format_banner_line(), +})) +""" + project_root = Path(__file__).resolve().parents[1] + env = { + **os.environ, + "PYTHONPATH": os.pathsep.join( + filter(None, (str(project_root), os.environ.get("PYTHONPATH"))) + ), + "AMBIENT_MARKER": str(ambient_marker), + } + result = subprocess.run( + [ + sys.executable, + "-c", + script, + str(ambient_root), + "entry" if source_on_entry else "external", + str(source_root), + ], + text=True, + capture_output=True, + env=env, + timeout=30, + ) + assert result.returncode == 0, result.stderr + result.stdout + assert not ambient_marker.exists() + outcome = json.loads(result.stdout) + assert outcome["provider_name"] == "Example" + assert outcome["banner"] == ( + "Bundle: fixture | Provider: Example | fixture-model" + ) + + +@pytest.mark.parametrize( + ("required", "value", "expect_error"), + [ + (True, "configured", False), + (False, "${EXAMPLE_WORK_API_KEY}", False), + (True, "${EXAMPLE_WORK_API_KEY}", True), + ], + ids=["set-credential", "optional-unset", "required-unset"], +) +@pytest.mark.parametrize( + "source_location", + ["entry", "modules", "override"], + ids=["entry", "modules", "override"], +) +def test_fresh_process_source_preflight_uses_configured_provider_root( + tmp_path, required, value, expect_error, source_location +): + """Exercise CLI preflight, Foundation resolver, and core's real module path. + + The process has both an activated configured source and an installed-looking + conflicting entry point. Only the configured source may provide metadata + or mount the provider. + """ + source_root = tmp_path / "configured-source" + ambient_root = tmp_path / "ambient" + source_marker = tmp_path / "source-marker" + ambient_marker = tmp_path / "ambient-marker" + _write_standard_provider(source_root, required=required) + _write_ambient_entry_point(ambient_root) + assert (ambient_root / "ambient_example-0.0.dist-info" / "METADATA").is_file() + assert ( + ambient_root / "ambient_example-0.0.dist-info" / "entry_points.txt" + ).is_file() + + script = r""" +import asyncio +import importlib +import json +import os +import sys +from pathlib import Path +from types import SimpleNamespace + +from amplifier_foundation.bundle import Bundle, BundleModuleResolver, PreparedBundle +import amplifier_app_cli.lib.bundle_loader as bundle_loader +import amplifier_app_cli.lib.bundle_loader.prepare as prepare +import amplifier_app_cli.paths as paths +from amplifier_app_cli.runtime.config import resolve_bundle_config + +source_root = Path(sys.argv[1]) +ambient_root = Path(sys.argv[2]) +value = sys.argv[3] +source_location = sys.argv[4] +sys.path.insert(0, str(ambient_root)) +# This mirrors Foundation's activator contract: selected roots are already +# active ahead of installed packages when the prepared resolver is used. +sys.path.insert(0, str(source_root)) + +entry = { + "module": "provider-example", + "id": "example-work", + "config": {"api_key": value}, +} +if source_location == "entry": + entry["source"] = str(source_root) +prepared = PreparedBundle( + mount_plan={ + "session": { + "orchestrator": "orchestrator-example", + "context": "context-example", + }, + "providers": [entry], + }, + resolver=BundleModuleResolver({ + "provider-example": source_root, + "orchestrator-example": source_root, + "context-example": source_root, + }), + bundle=Bundle(name="test", base_path=source_root), +) + +async def fake_load_and_prepare_bundle(*args, **kwargs): + expected_sources = ( + {} if source_location == "entry" else {"provider-example": str(source_root)} + ) + assert kwargs.get("source_overrides") == expected_sources or ( + expected_sources == {} and kwargs.get("source_overrides") is None + ) + return prepared + +prepare.load_and_prepare_bundle = fake_load_and_prepare_bundle +bundle_loader.AppBundleDiscovery = lambda **kwargs: object() +paths.get_bundle_search_paths = lambda: [] + +class Settings: + def get_app_bundles(self): return [] + def get_source_overrides(self): + return {"provider-example": str(source_root)} if source_location == "override" else {} + def get_module_sources(self): + return {"provider-example": str(source_root)} if source_location == "modules" else {} + def get_bundle_sources(self): return {} + def get_provider_overrides(self): return [] + def get_config_overrides(self): return {} + def get_tool_overrides(self, **kwargs): return [] + def get_notification_hook_overrides(self): return [] + def get_routing_config(self): return None + def get_notification_flags(self): + return SimpleNamespace(desktop_enabled=False, push_enabled=False) + +async def main(): + try: + _, configured = await resolve_bundle_config("test", Settings()) + except ValueError as exc: + print(json.dumps({"error": str(exc)})) + return + session = await configured.create_session() + try: + provider = session.coordinator.get("providers")["example-work"] + module = importlib.import_module("amplifier_module_provider_example") + print(json.dumps({ + "origin": str(Path(module.__file__).resolve()), + "provider_name": provider.name, + })) + finally: + await session.cleanup() + +asyncio.run(main()) +""" + project_root = Path(__file__).resolve().parents[1] + env = { + **os.environ, + "PYTHONPATH": os.pathsep.join( + filter(None, (str(project_root), os.environ.get("PYTHONPATH"))) + ), + "SOURCE_MARKER": str(source_marker), + "AMBIENT_MARKER": str(ambient_marker), + "EXAMPLE_WORK_API_KEY": "", + } + result = subprocess.run( + [ + sys.executable, + "-c", + script, + str(source_root), + str(ambient_root), + value, + source_location, + ], + check=False, + text=True, + capture_output=True, + env=env, + timeout=30, + ) + assert result.returncode == 0, ( + f"child failed ({result.returncode})\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) + outcome = json.loads(result.stdout) + + assert not ambient_marker.exists() + if expect_error: + assert "EXAMPLE_WORK_API_KEY" in outcome["error"] + assert not source_marker.exists() + else: + assert outcome["provider_name"] == "example" + assert Path(outcome["origin"]) == ( + source_root / "amplifier_module_provider_example" / "__init__.py" + ) + assert source_marker.read_text(encoding="utf-8") == "mounted" + + +def test_source_path_rejects_cached_wrong_root_without_entry_point_load( + monkeypatch, tmp_path +): + """A cached same-name package from another root cannot satisfy a source pin.""" + from amplifier_app_cli.provider_loader import get_provider_info + + source_root = tmp_path / "configured" + ambient_root = tmp_path / "ambient" + _write_standard_provider(source_root, required=True) + _write_ambient_entry_point(ambient_root) + module_name = "amplifier_module_provider_example" + cached_module = type(sys)(module_name) + cached_module.__file__ = str(ambient_root / module_name / "__init__.py") + monkeypatch.setitem(sys.modules, module_name, cached_module) + + with patch( + "amplifier_app_cli.provider_loader.importlib.metadata.entry_points" + ) as entry_points, pytest.raises(ImportError, match="already loaded"): + get_provider_info("provider-example", source_path=source_root) + + assert sys.modules[module_name] is cached_module + entry_points.assert_not_called() + + +def test_source_path_missing_package_never_falls_back_to_ambient( + monkeypatch, tmp_path +): + """An incomplete selected root fails before discovery can import ambient code.""" + from amplifier_app_cli.provider_loader import get_provider_info + + source_root = tmp_path / "configured" + source_root.mkdir() + ambient_root = tmp_path / "ambient" + ambient_marker = tmp_path / "ambient-marker" + _write_ambient_entry_point(ambient_root) + monkeypatch.setenv("AMBIENT_MARKER", str(ambient_marker)) + monkeypatch.syspath_prepend(str(ambient_root)) + + with patch( + "amplifier_app_cli.provider_loader.importlib.metadata.entry_points" + ) as entry_points, pytest.raises(ImportError, match="does not contain"): + get_provider_info("provider-example", source_path=source_root) + + assert not ambient_marker.exists() + entry_points.assert_not_called()