From ab631f10618d5b98e7807c54ba94c8ce6c8c38f8 Mon Sep 17 00:00:00 2001 From: Brian Krabach Date: Wed, 23 Sep 2026 23:31:30 -0700 Subject: [PATCH] Expose provider preparation and pre-initialization host policy --- amplifier_foundation/bundle/__init__.py | 3 + amplifier_foundation/bundle/_dataclass.py | 62 +++- amplifier_foundation/bundle/_prepared.py | 57 +++- .../bundle/_provider_preparation.py | 83 +++++ docs/provider-failure-policy.md | 61 ++++ tests/test_provider_preparation_policy.py | 298 ++++++++++++++++++ 6 files changed, 557 insertions(+), 7 deletions(-) create mode 100644 amplifier_foundation/bundle/_provider_preparation.py create mode 100644 docs/provider-failure-policy.md create mode 100644 tests/test_provider_preparation_policy.py diff --git a/amplifier_foundation/bundle/__init__.py b/amplifier_foundation/bundle/__init__.py index d72f6fa6..e9566d05 100644 --- a/amplifier_foundation/bundle/__init__.py +++ b/amplifier_foundation/bundle/__init__.py @@ -9,7 +9,10 @@ PreparedBundle, ) +from amplifier_foundation.bundle._provider_preparation import ProviderPreparationFailure + __all__ = [ + "ProviderPreparationFailure", "Bundle", "BundleModuleResolver", "BundleModuleSource", diff --git a/amplifier_foundation/bundle/_dataclass.py b/amplifier_foundation/bundle/_dataclass.py index 7ab49fec..293fab54 100644 --- a/amplifier_foundation/bundle/_dataclass.py +++ b/amplifier_foundation/bundle/_dataclass.py @@ -11,6 +11,7 @@ if TYPE_CHECKING: from amplifier_foundation.bundle._prepared import PreparedBundle + from amplifier_foundation.bundle._provider_preparation import ProviderFailurePolicy from amplifier_foundation.bundle._provenance import ( _prov_add as _prov_add, # re-exported for backwards compatibility @@ -335,6 +336,7 @@ async def prepare( cache_dir: Path | None = None, refresh_dependencies: bool = False, install_overrides: Path | None = None, + provider_failure_policy: ProviderFailurePolicy | None = None, ) -> PreparedBundle: """Prepare bundle for execution by activating all modules. @@ -372,6 +374,12 @@ async def prepare( install_overrides: Explicit uv overrides file for host-qualified dependencies such as native wheels. Replaces automatic overrides based on installed versions; the file is never modified. + provider_failure_policy: Optional async callback receiving a private + ProviderPreparationFailure. Returning accepts that provider's + source-resolution/activation failure; raising aborts preparation. + Provider entries remain in the plan, with failed sources blocked + in the prepared resolver. Non-provider strictness and bundle + package failures are unchanged. Not a routing fallback policy. Returns: PreparedBundle with mount_plan and create_session() helper. @@ -415,7 +423,17 @@ def resolve_with_overrides(module_id: str, source: str) -> str: install_overrides=install_overrides, ) - # Collect all modules that need activation + # Only explicit host policy opts providers out of strict aggregate + # preparation. Other module and bundle-package rules stay intact. + from amplifier_foundation.bundle._provider_preparation import ( + ProviderPreparation, + ) + + provider_preparation = ( + ProviderPreparation(provider_failure_policy) + if provider_failure_policy + else None + ) modules_to_activate = [] # Helper to apply source resolver if provided @@ -442,7 +460,10 @@ def resolve_source(mod_spec: dict) -> dict: for section in ["providers", "tools", "hooks"]: for mod_spec in mount_plan.get(section, []): if isinstance(mod_spec, dict) and "source" in mod_spec: - modules_to_activate.append(resolve_source(mod_spec)) + if section == "providers" and provider_preparation is not None: + provider_preparation.specs.append((deepcopy(mod_spec), None)) + else: + modules_to_activate.append(resolve_source(mod_spec)) # Pre-activate modules declared in agent configs so child sessions # can find them via the inherited BundleModuleResolver. @@ -470,7 +491,21 @@ def resolve_source(mod_spec: dict) -> dict: if isinstance(agent_mods, list): for mod_spec in agent_mods: if isinstance(mod_spec, dict) and "source" in mod_spec: - modules_to_activate.append(resolve_source(mod_spec)) + if ( + agent_section == "providers" + and provider_preparation is not None + ): + provider_preparation.specs.append( + (deepcopy(mod_spec), _agent_name) + ) + else: + modules_to_activate.append(resolve_source(mod_spec)) + + provider_sources = ( + await provider_preparation.resolve(resolve_source) + if provider_preparation is not None + else [] + ) # Phase 1: schema-only modes walk. # Validates contributes structure; actual module activation for contributed @@ -489,7 +524,7 @@ def resolve_source(mod_spec: dict) -> dict: if install_deps: declared_sources = [ m["source"] - for m in modules_to_activate + for m in [*modules_to_activate, *provider_sources] if isinstance(m.get("source"), str) ] # This bundle's own package: a failure here is a failure of the bundle @@ -525,12 +560,26 @@ def resolve_source(mod_spec: dict) -> dict: modules_to_activate, progress_callback=progress_callback ) + if provider_preparation is not None: + module_paths.update( + await provider_preparation.activate(activator, progress_callback) + ) + # Save install state to disk for fast subsequent startups activator.finalize() # Create resolver from activated paths with activator for lazy activation # This enables child sessions to activate agent-specific modules on-demand - resolver = BundleModuleResolver(module_paths, activator=activator) + resolver = BundleModuleResolver( + module_paths, + activator=activator, + source_paths=provider_preparation.source_paths + if provider_preparation + else None, + source_failures=provider_preparation.failed_sources + if provider_preparation + else None, + ) # Get bundle package paths for inheritance by child sessions bundle_package_paths = activator.bundle_package_paths @@ -548,6 +597,9 @@ def resolve_source(mod_spec: dict) -> dict: bundle_package_paths=bundle_package_paths, mode_warnings=mode_warnings, module_exports=module_exports, + provider_preparation_failures=tuple(provider_preparation.failures) + if provider_preparation + else (), ) def resolve_context_path(self, name: str) -> Path | None: diff --git a/amplifier_foundation/bundle/_prepared.py b/amplifier_foundation/bundle/_prepared.py index feade023..04fee88a 100644 --- a/amplifier_foundation/bundle/_prepared.py +++ b/amplifier_foundation/bundle/_prepared.py @@ -24,6 +24,7 @@ from amplifier_foundation.spawn_utils import apply_provider_preferences_with_resolution from amplifier_foundation.bundle._dataclass import Bundle +from amplifier_foundation.bundle._provider_preparation import ProviderPreparationFailure logger = logging.getLogger(__name__) @@ -190,6 +191,9 @@ def __init__( self, module_paths: dict[str, Path], activator: "ModuleActivator | None" = None, + *, + source_paths: dict[tuple[str, str], Path] | None = None, + source_failures: dict[tuple[str, str], Exception] | None = None, ) -> None: """Initialize with activated module paths and optional activator. @@ -200,8 +204,38 @@ def __init__( """ self._paths = module_paths self._activator = activator + self._source_paths = dict(source_paths or {}) + self._source_failures = dict(source_failures or {}) + self._prepared_provider_ids = { + key[0] for key in (*self._source_paths, *self._source_failures) + } self._activation_lock = asyncio.Lock() + def _source_path(self, module_id: str, hint: Any) -> Path | None: + if module_id not in self._prepared_provider_ids: + return None + if not isinstance(hint, str): + raise CoreModuleNotFoundError( + f"An explicit prepared source is required for provider '{module_id}'" + ) + key = (module_id, hint) + if key in self._source_failures: + # Do not retry an accepted failed source or reuse another account's + # source under the same module ID during this prepared generation. + raise CoreModuleNotFoundError( + f"Provider source preparation failed for '{module_id}'" + ) from self._source_failures[key] + if key in self._source_paths: + path = self._source_paths[key] + if not path.exists(): + raise CoreModuleNotFoundError( + f"Prepared provider source is no longer available for '{module_id}'" + ) + return path + raise CoreModuleNotFoundError( + f"Provider source was not prepared for '{module_id}'; prepare the updated bundle first" + ) + def resolve( self, module_id: str, source_hint: Any = None, profile_hint: Any = None ) -> BundleModuleSource: @@ -220,8 +254,8 @@ def resolve( FIXME: Remove profile_hint parameter after all callers migrate to source_hint (target: v2.0). """ - _hint = profile_hint if profile_hint is not None else source_hint # noqa: F841 - path = self._paths.get(module_id) + hint = profile_hint if profile_hint is not None else source_hint + path = self._source_path(module_id, hint) or self._paths.get(module_id) if path is None or not path.exists(): raise CoreModuleNotFoundError( f"Module '{module_id}' has no available path in prepared bundle. " @@ -249,6 +283,9 @@ async def async_resolve( FIXME: Remove profile_hint parameter after all callers migrate to source_hint (target: v2.0). """ hint = profile_hint if profile_hint is not None else source_hint + source_path = self._source_path(module_id, hint) + if source_path is not None: + return BundleModuleSource(source_path) # Prepared bundles can outlive a cache checkout (for example when a # child is spawned after cache eviction). Reuse only a present path; # let the activator recover a missing checkout from the declared source. @@ -331,6 +368,7 @@ class PreparedBundle: bundle_package_paths: list[str] = field(default_factory=list) mode_warnings: list[str] = field(default_factory=list) module_exports: dict[str, list[str]] = field(default_factory=dict) + provider_preparation_failures: tuple[ProviderPreparationFailure, ...] = () def _build_bundles_for_resolver(self, bundle: "Bundle") -> dict[str, "Bundle"]: """Build bundle registry for mention resolution. @@ -546,6 +584,8 @@ async def create_session( display_system: Any = None, session_cwd: Path | None = None, is_resumed: bool = False, + *, + before_initialize: Callable[[Any], Awaitable[None]] | None = None, ) -> Any: """Create an AmplifierSession with the resolver properly mounted. @@ -569,6 +609,10 @@ async def create_session( Defaults to bundle.base_path if not provided. is_resumed: Whether this session is being resumed (vs newly created). Controls whether session:start or session:resume events are emitted. + before_initialize: Optional async host callback receiving the newly + created session after resolver/working-directory capabilities are + mounted, before module initialization and lifecycle hooks. May + install failure policy or raise to abort. Not inherited implicitly. Returns: Initialized AmplifierSession ready for execute(). @@ -657,6 +701,9 @@ async def create_session( "mention_deduplicator", initial_deduplicator ) + # Host policy must be registered before mounts and lifecycle routing. + if before_initialize is not None: + await before_initialize(session) # Initialize the session (loads all modules) await session.initialize() @@ -754,6 +801,7 @@ async def spawn( session_cwd: Path | None = None, provider_preferences: list[ProviderPreference] | None = None, self_delegation_depth: int = 0, + before_initialize: Callable[[Any], Awaitable[None]] | None = None, ) -> dict[str, Any]: """Spawn a sub-session with a child bundle. @@ -773,6 +821,9 @@ async def spawn( Args: child_bundle: Bundle to spawn (already resolved by app layer). instruction: Task instruction for the sub-session. + before_initialize: Optional async host callback for the fresh child, + before its modules and lifecycle hooks initialize. Hosts should + pass the same policy installer used for root/resumed sessions. compose: Whether to compose child with parent bundle (default True). parent_session: Parent session for lineage tracking and UX inheritance. session_id: Optional session ID for resuming existing session. @@ -944,6 +995,8 @@ async def spawn( "mention_deduplicator", ContentDeduplicator() ) + if before_initialize is not None: + await before_initialize(child_session) await child_session.initialize() # Register mentions:resolved on observability.events for child sessions. diff --git a/amplifier_foundation/bundle/_provider_preparation.py b/amplifier_foundation/bundle/_provider_preparation.py new file mode 100644 index 00000000..c995a6e3 --- /dev/null +++ b/amplifier_foundation/bundle/_provider_preparation.py @@ -0,0 +1,83 @@ +"""Opt-in provider preparation outcomes for application-owned failure policy.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable +from copy import deepcopy +from dataclasses import dataclass +from pathlib import Path +from typing import Any + + +@dataclass(frozen=True) +class ProviderPreparationFailure: + """Private diagnostic input. The spec and exception may contain credentials.""" + + provider_spec: dict[str, Any] + agent_name: str | None + phase: str + error: Exception + + +ProviderFailurePolicy = Callable[[ProviderPreparationFailure], Awaitable[None]] + + +class ProviderPreparation: + """Retain source identity while handling only explicitly opted-in providers.""" + + def __init__(self, policy: ProviderFailurePolicy): + self.policy = policy + self.specs: list[tuple[dict, str | None]] = [] + self.resolved: list[tuple[dict, dict, str | None]] = [] + self.failures: list[ProviderPreparationFailure] = [] + self.failed_sources: dict[tuple[str, str], Exception] = {} + self.source_paths: dict[tuple[str, str], Path] = {} + + async def failed(self, spec, agent_name, phase, error, resolved=None): + failure = ProviderPreparationFailure(deepcopy(spec), agent_name, phase, error) + # Raising from host policy aborts; returning accepts this preparation + # failure, without silently removing the configured entry. + await self.policy( + ProviderPreparationFailure(deepcopy(spec), agent_name, phase, error) + ) + self.failures.append(failure) + for source in (spec, resolved or spec): + self.failed_sources[(source["module"], source["source"])] = error + + async def resolve(self, source_resolver): + for spec, agent_name in self.specs: + if not spec.get("module") or not spec.get("source"): + continue # Match activate_all's handling of incomplete rows. + try: + resolved = source_resolver(deepcopy(spec)) + except Exception as error: # noqa: BLE001 - host policy owns module failures + await self.failed(spec, agent_name, "source_resolution", error) + else: + self.resolved.append((spec, resolved, agent_name)) + return [resolved for _, resolved, _ in self.resolved] + + async def activate(self, activator, progress_callback=None): + keys = list( + dict.fromkeys((row["module"], row["source"]) for _, row, _ in self.resolved) + ) + results = await asyncio.gather( + *[ + activator.activate(module, source, progress_callback=progress_callback) + for module, source in keys + ], + return_exceptions=True, + ) + by_source = dict(zip(keys, results)) + paths = {} + for spec, resolved, agent_name in self.resolved: + result = by_source[(resolved["module"], resolved["source"])] + if isinstance(result, Exception): + await self.failed(spec, agent_name, "activation", result, resolved) + elif isinstance(result, BaseException): + raise result # Never accept cancellation or process termination. + else: + paths[spec["module"]] = result + for source in (spec, resolved): + self.source_paths[(source["module"], source["source"])] = result + return paths diff --git a/docs/provider-failure-policy.md b/docs/provider-failure-policy.md new file mode 100644 index 00000000..af1b75a1 --- /dev/null +++ b/docs/provider-failure-policy.md @@ -0,0 +1,61 @@ +# Provider preparation and initialization policy + +Applications can opt into per-provider source failure handling without disabling +strict preparation for tools, hooks, orchestrators, contexts, or bundle packages. +The default `Bundle.prepare()` behavior is unchanged. + +```python +from amplifier_foundation.bundle import ProviderPreparationFailure + +async def record_failure(outcome: ProviderPreparationFailure): + # Private host-owned diagnostics, not public browser state. + failures.append(outcome) + # Returning accepts this provider failure; raising rejects preparation. + +prepared = await bundle.prepare( + strict=True, + provider_failure_policy=record_failure, +) + +async def install_host_policy(session): + session.coordinator.register_capability( + "provider.load_failure", handle_mount_failure + ) + +session = await prepared.create_session(before_initialize=install_host_policy) +# Explicitly apply the same host policy to spawned sessions: +result = await prepared.spawn( + child_bundle, instruction, before_initialize=install_host_policy +) +``` + +`ProviderPreparationFailure` contains a copy of the configured provider entry, +its agent name (`None` for root entries), a phase (`source_resolution` or +`activation`), and the original exception. Accepted outcomes are also available +in `prepared.provider_preparation_failures`. Specs and exceptions may include +credentials; applications must expose only their own allowlisted diagnostics. +Changing the callback's spec does not modify the bundle, mount plan, or stored +outcome. Cancellation is propagated, never recorded as an accepted failure. + +Provider entries, instance IDs, source declarations, model configuration, and +priority remain in the mount plan. The resolver records the exact failed source; +it cannot borrow a successful path for another source using the same module ID, +or silently retry an accepted failure. Successful source aliases are also +recorded explicitly. Repeated accounts using one source share its activation, +while each failed account gets its own outcome. An unknown source, missing +prepared path, or omitted source for these opted-in providers fails explicitly; +prepare the updated bundle again to recover or introduce a new provider source. +Ordinary non-provider lazy resolution is unchanged. + +`before_initialize` is awaited after the module resolver, working directory, and +mention capabilities are mounted, but before module initialization and lifecycle +hooks. It applies to root, resumed, and spawned sessions and can deliberately +abort. It is not inherited automatically: the application owns spawning and must +pass its installer for each new session. This general callback does not require +a particular provider failure capability; applications using +`provider.load_failure` need a Core version that supplies that contract. + +These mechanisms do not choose providers, remove unavailable entries, or enable +account fallback. The host still needs an unavailable-provider representation and +routing that distinguishes an explicit failed account from an absent model match. +They do not make arbitrary package-install side effects transactional. diff --git a/tests/test_provider_preparation_policy.py b/tests/test_provider_preparation_policy.py new file mode 100644 index 00000000..0ef5ae0d --- /dev/null +++ b/tests/test_provider_preparation_policy.py @@ -0,0 +1,298 @@ +"""Opt-in preparation policy preserves identity, strictness, and lifecycle order.""" + +import asyncio +from copy import deepcopy +from unittest.mock import AsyncMock + +import pytest +from amplifier_core import AmplifierSession +from amplifier_core.module_sources import ModuleNotFoundError + +from amplifier_foundation.bundle import Bundle, BundleModuleResolver, PreparedBundle +from amplifier_foundation.modules.activator import ModuleActivator + + +@pytest.mark.asyncio +async def test_failed_provider_source_does_not_replace_healthy_same_module( + tmp_path, monkeypatch +): + calls = [] + + async def activate(self, module, source, **kwargs): + calls.append((module, source)) + if source == "missing": + raise OSError("private source error") + p = tmp_path / source + p.mkdir(exist_ok=True) + return p + + monkeypatch.setattr(ModuleActivator, "activate", activate) + specs = [ + { + "module": "provider-x", + "instance_id": "bad", + "source": "missing", + "config": {"priority": 1}, + }, + { + "module": "provider-x", + "instance_id": "good", + "source": "healthy", + "config": {"priority": 2}, + }, + {"module": "provider-x", "instance_id": "another", "source": "healthy"}, + ] + b = Bundle( + name="test", + providers=specs, + agents={ + "child": { + "providers": [ + { + "module": "provider-x", + "instance_id": "child-account", + "source": "missing", + } + ] + } + }, + ) + original = deepcopy(b.to_mount_plan()) + failures = [] + + async def accept(failure): + failures.append(failure) + # Host callbacks own a copy; cannot rewrite the plan or stored outcome. + failure.provider_spec["config"] = {"changed": True} + + prepared = await b.prepare( + install_deps=False, + cache_dir=tmp_path / "cache", + strict=True, + provider_failure_policy=accept, + ) + assert prepared.mount_plan == original == b.to_mount_plan() + assert len(failures) == 2 + assert [f.agent_name for f in failures] == [None, "child"] + assert prepared.provider_preparation_failures[0].provider_spec == specs[0] + assert calls.count(("provider-x", "healthy")) == 1 + assert calls.count(("provider-x", "missing")) == 1 + assert ( + await prepared.resolver.async_resolve("provider-x", "healthy") + ).resolve() == tmp_path / "healthy" + for source in ("missing", "different", None): + with pytest.raises(ModuleNotFoundError): + await prepared.resolver.async_resolve("provider-x", source) + with pytest.raises(ModuleNotFoundError): + prepared.resolver.resolve("provider-x", source) + assert len(calls) == 2 # No lazy retry or other-source substitution. + + +@pytest.mark.asyncio +async def test_source_resolution_failure_retains_original_identity( + tmp_path, monkeypatch +): + async def activate(self, module, source, **kwargs): + return tmp_path + + monkeypatch.setattr(ModuleActivator, "activate", activate) + + def resolve(module, source): + if source == "bad": + raise ValueError("private override failed") + return "resolved-good" + + failures = [] + + async def accept(outcome): + failures.append(outcome) + + b = Bundle( + name="test", + providers=[ + {"module": "provider-x", "source": "bad", "instance_id": "broken"}, + {"module": "provider-y", "source": "good"}, + ], + ) + p = await b.prepare( + install_deps=False, + cache_dir=tmp_path, + strict=True, + source_resolver=resolve, + provider_failure_policy=accept, + ) + assert failures[0].phase == "source_resolution" + assert failures[0].provider_spec["source"] == "bad" + assert (await p.resolver.async_resolve("provider-y", "good")).resolve() == tmp_path + assert ( + await p.resolver.async_resolve("provider-y", "resolved-good") + ).resolve() == tmp_path + with pytest.raises(ModuleNotFoundError): + await p.resolver.async_resolve("provider-x", "bad") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("section", ["tools", "hooks", "providers"]) +async def test_default_strictness_and_non_provider_errors_still_abort( + tmp_path, monkeypatch, section +): + async def fail(*args, **kwargs): + raise OSError("cannot activate") + + monkeypatch.setattr(ModuleActivator, "activate", fail) + policy = AsyncMock() + b = Bundle(name="test", **{section: [{"module": "anything", "source": "bad"}]}) + args = {"provider_failure_policy": policy} if section != "providers" else {} + with pytest.raises(Exception, match="strict mode"): + await b.prepare(install_deps=False, cache_dir=tmp_path, strict=True, **args) + policy.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_host_can_reject_preparation_failure(tmp_path, monkeypatch): + async def fail(*args, **kwargs): + raise OSError("failed source") + + async def reject(outcome): + raise RuntimeError("host rejects incomplete configuration") + + monkeypatch.setattr(ModuleActivator, "activate", fail) + b = Bundle(name="test", providers=[{"module": "provider-x", "source": "bad"}]) + with pytest.raises(RuntimeError, match="host rejects"): + await b.prepare( + install_deps=False, cache_dir=tmp_path, provider_failure_policy=reject + ) + + +@pytest.mark.asyncio +async def test_cancellation_is_not_accepted_as_provider_failure(tmp_path, monkeypatch): + async def cancel(*args, **kwargs): + raise asyncio.CancelledError() + + monkeypatch.setattr(ModuleActivator, "activate", cancel) + policy = AsyncMock() + b = Bundle(name="test", providers=[{"module": "provider-x", "source": "bad"}]) + with pytest.raises(asyncio.CancelledError): + await b.prepare( + install_deps=False, cache_dir=tmp_path, provider_failure_policy=policy + ) + policy.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("resumed", [False, True]) +async def test_host_callback_precedes_root_and_resume_initialization( + tmp_path, monkeypatch, resumed +): + order = [] + + async def initialize(session): + order.append("initialize") + assert ( + session.coordinator.get_capability("provider.load_failure") == "host-marker" + ) + assert session.coordinator.get_capability("mention_resolver") is not None + + monkeypatch.setattr(AmplifierSession, "initialize", initialize) + b = Bundle(name="test", session={"orchestrator": "test", "context": "test"}) + p = PreparedBundle( + mount_plan=b.to_mount_plan(), bundle=b, resolver=BundleModuleResolver({}) + ) + + async def before(session): + order.append("host") + assert session.coordinator.get_capability("session.working_dir") == str( + tmp_path + ) + session.coordinator.register_capability("provider.load_failure", "host-marker") + + await p.create_session( + is_resumed=resumed, session_cwd=tmp_path, before_initialize=before + ) + assert order == ["host", "initialize"] + + +@pytest.mark.asyncio +async def test_host_callback_failure_prevents_module_initialization( + tmp_path, monkeypatch +): + initialize = AsyncMock() + monkeypatch.setattr(AmplifierSession, "initialize", initialize) + b = Bundle(name="test", session={"orchestrator": "test", "context": "test"}) + p = PreparedBundle( + mount_plan=b.to_mount_plan(), bundle=b, resolver=BundleModuleResolver({}) + ) + + async def before(session): + raise ValueError("host validation failed") + + with pytest.raises(ValueError, match="host validation"): + await p.create_session(before_initialize=before) + initialize.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_child_callback_runs_before_mount_and_preserves_bundle( + tmp_path, monkeypatch +): + observed = [] + + async def initialize(session): + assert ( + session.coordinator.get_capability("provider.load_failure") + == "child-policy" + ) + observed.append(("initialized", session.session_id)) + + async def execute(session, instruction): + assert instruction == "synthetic instruction" + return "done" + + monkeypatch.setattr(AmplifierSession, "initialize", initialize) + monkeypatch.setattr(AmplifierSession, "execute", execute) + b = Bundle(name="test", session={"orchestrator": "test", "context": "test"}) + original = deepcopy(b.to_mount_plan()) + p = PreparedBundle( + mount_plan=b.to_mount_plan(), bundle=b, resolver=BundleModuleResolver({}) + ) + + async def before(session): + assert session.coordinator.get_capability("mention_resolver") is not None + assert session.coordinator.get_capability("session.working_dir") == str( + tmp_path + ) + session.coordinator.register_capability("provider.load_failure", "child-policy") + observed.append(("host", session.session_id)) + + result = await p.spawn( + Bundle(name="child"), + "synthetic instruction", + session_cwd=tmp_path, + before_initialize=before, + ) + assert [row[0] for row in observed] == ["host", "initialized"] + assert observed[0][1] == result["session_id"] + assert b.to_mount_plan() == original + + +@pytest.mark.asyncio +async def test_package_failure_is_not_a_provider_failure(tmp_path, monkeypatch): + from amplifier_foundation.modules.activator import BundlePackageInstallError + + async def fail(*args, **kwargs): + raise BundlePackageInstallError( + tmp_path, "synthetic-package", "package preparation failed" + ) + + monkeypatch.setattr(ModuleActivator, "activate_bundle_package", fail) + policy = AsyncMock() + b = Bundle( + name="test", + base_path=tmp_path, + providers=[{"module": "provider-x", "source": "./modules/provider-x"}], + ) + with pytest.raises(BundlePackageInstallError): + await b.prepare( + cache_dir=tmp_path / "cache", strict=True, provider_failure_policy=policy + ) + policy.assert_not_awaited()