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
3 changes: 3 additions & 0 deletions amplifier_foundation/bundle/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@
PreparedBundle,
)

from amplifier_foundation.bundle._provider_preparation import ProviderPreparationFailure

__all__ = [
"ProviderPreparationFailure",
"Bundle",
"BundleModuleResolver",
"BundleModuleSource",
Expand Down
62 changes: 57 additions & 5 deletions amplifier_foundation/bundle/_dataclass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down
57 changes: 55 additions & 2 deletions amplifier_foundation/bundle/_prepared.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

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

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

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

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

Expand All @@ -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.
Expand Down Expand Up @@ -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.
Expand Down
83 changes: 83 additions & 0 deletions amplifier_foundation/bundle/_provider_preparation.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading