Skip to content
Merged
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
18 changes: 4 additions & 14 deletions amplifier_app_cli/effective_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
66 changes: 61 additions & 5 deletions amplifier_app_cli/provider_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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.
Expand All @@ -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)
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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}'"
Expand All @@ -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}"
)
Expand Down
57 changes: 49 additions & 8 deletions amplifier_app_cli/runtime/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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):
Expand All @@ -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

Expand Down
Loading
Loading