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
25 changes: 25 additions & 0 deletions docs/CAPABILITY_REGISTRY.md
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,31 @@ else:
|------------|----------|----------|----------|
| `session.spawn` | `async (agent_name: str, task: str, parent_session) → dict` | amplifier-app-cli | tool-task |
| `session.resume` | `async (session_id: str, task: str) → dict` | amplifier-app-cli | tool-task |
| `provider.load_failure` | `async (coordinator, provider_spec: dict, error: Exception) → None` | host application | Python session initialization |

### Provider initialization failures

Register `provider.load_failure` before session initialization to apply host
policy when a configured provider cannot load, mount, or finish instance remapping.
It runs once for the failed configuration entry, after restoring the provider
mount table and emitting `module:load_failed`, and before loading the next entry.
The spec is a deep copy of that entry, including its source, `instance_id`, and
configuration. The callback may mount a host-owned unavailable-provider marker,
record an outcome, or raise to abort initialization. Callback exceptions propagate;
observability event-handler exceptions do not.

Without a callback, provider failure keeps the existing warn-and-continue policy.
Core does not select a replacement, prune configuration, reorder providers, or
decide whether a failed provider is required. Hosts must preserve configured
identity and implement routing policy themselves. A callback that retains a
failure marker should use the configured instance identity and a safe reason code;
the spec and exception may contain secrets and must not be copied to public state.

Restoration covers the provider mount table, including an overwritten default
slot and mounts added by a failed attempt. It does not undo arbitrary external
side effects inside third-party mount functions. A failed instance's readiness
callback is not queued. Rollback failure aborts initialization rather than asking
the host to recover against a partially restored table.

## Pattern

Expand Down
6 changes: 6 additions & 0 deletions docs/module-failure-reasons.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,4 +6,10 @@ Categories are `invalid_package_layout`, `missing_source`, `invalid_entry_point`

The legacy `error` field remains unchanged for compatibility and is not a safe public diagnostic. Consumers should allowlist `reason_code` and supply their own remediation text. This addition changes neither the provider/tool/hook nonfatal session policy nor host decisions about incomplete configured sessions.

Provider failures also include `instance_id` when one was explicitly configured.
This distinguishes failed accounts using the same provider module. The optional
[`provider.load_failure`](CAPABILITY_REGISTRY.md#provider-initialization-failures)
capability lets hosts apply policy before lifecycle routing; the event itself
remains an observability mechanism.

Validation used Python source over the installed native kernel: 44 loader and session-initialization tests passed, including fixtures for the four requested categories, unknown and forged codes, and event compatibility. No native code or provider protocol changed.
45 changes: 36 additions & 9 deletions python/amplifier_core/_session_init.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
"""

import logging
from copy import deepcopy
from typing import Any

logger = logging.getLogger(__name__)
Expand All @@ -20,16 +21,20 @@ def _safe_exception_str(e: BaseException) -> str:


async def _emit_module_load_failed(
coordinator: Any, module_type: str, module_id: str, error: BaseException
coordinator: Any,
module_type: str,
module_id: str,
error: BaseException,
*,
instance_id: str | None = None,
) -> None:
"""Emit the module:load_failed observability event for a provider, tool,
or hook that raised during load/mount.

This is a mechanism only: the kernel makes the failure observable via the
canonical event stream. It does not decide whether the session should
abort -- that policy choice belongs to a hook module subscribed to this
event. Mirrors the on_session_ready failure pattern below: event emission
failure must never suppress the original WARNING log.
canonical event stream. Event subscribers cannot abort initialization:
their failures are isolated from the original WARNING log. Provider abort
or recovery policy belongs to the optional provider.load_failure capability.
"""
from .events import MODULE_LOAD_FAILED
from .loader import module_failure_reason
Expand All @@ -42,6 +47,7 @@ async def _emit_module_load_failed(
"module_id": module_id,
"error": _safe_exception_str(error),
"reason_code": module_failure_reason(error),
**({"instance_id": instance_id} if instance_id else {}),
},
)
except Exception:
Expand Down Expand Up @@ -159,6 +165,7 @@ async def initialize_session(
if not module_id:
continue
instance_id = provider_config.get("instance_id") # multi-instance support
before_providers = dict(coordinator.get("providers") or {})
try:
logger.info(
f"Loading provider: {module_id}"
Expand Down Expand Up @@ -186,9 +193,6 @@ async def initialize_session(
cleanup = await provider_mount(coordinator)
if cleanup:
coordinator.register_cleanup(cleanup)
# B1 fix: enqueue on_session_ready ONLY after successful mount
if on_sr := getattr(provider_mount, "__on_session_ready__", None):
loader.enqueue_on_session_ready(on_sr[0], on_sr[1])

# Multi-instance remapping: if instance_id specified, remap mount name
if instance_id:
Expand All @@ -214,12 +218,35 @@ async def initialize_session(
logger.info(
f"Remapped provider '{default_name}' -> '{instance_id}'"
)
# Readiness belongs to a successfully mounted and remapped instance.
if on_sr := getattr(provider_mount, "__on_session_ready__", None):
loader.enqueue_on_session_ready(on_sr[0], on_sr[1])
except Exception as e:
logger.warning(
f"Failed to load provider '{module_id}': {_safe_exception_str(e)}",
exc_info=True,
)
await _emit_module_load_failed(coordinator, "provider", module_id, e)
# A provider may mount and then raise. Roll back this attempt before
# any host failure policy runs so an unrelated account is not lost.
for name in set(coordinator.get("providers") or {}) - set(before_providers):
await coordinator.unmount("providers", name=name)
# A failed attempt can remove a previous entry as well as overwrite
# it. Rebuild only in that case so default iteration order survives.
if list(coordinator.get("providers") or {}) != list(before_providers):
for name in list(coordinator.get("providers") or {}):
await coordinator.unmount("providers", name=name)
for name, previous in before_providers.items():
if (coordinator.get("providers") or {}).get(name) is not previous:
await coordinator.mount("providers", previous, name=name)
await _emit_module_load_failed(
coordinator, "provider", module_id, e, instance_id=instance_id
)
# Optional app policy, installed before initialization. Unlike an
# observability subscriber, this callback can deliberately abort.
# The kernel does not select another provider or change config.
handler = coordinator.get_capability("provider.load_failure")
if handler is not None:
await handler(coordinator, deepcopy(provider_config), e)

# Load tools
for tool_config in config.get("tools", []):
Expand Down
220 changes: 220 additions & 0 deletions tests/test_provider_failure_policy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,220 @@
"""Provider mount failures preserve sibling accounts before host policy runs."""

import copy
from unittest.mock import AsyncMock, MagicMock

import pytest
from amplifier_core import AmplifierSession
from amplifier_core._session_init import initialize_session
from amplifier_core.loader import ModuleLoader
from amplifier_core.testing import MockCoordinator


@pytest.fixture(params=["mock", "native"])
def coordinator(request):
if request.param == "mock":
return MockCoordinator()
# Exercise the installed native mount table and capability registry too.
return AmplifierSession(
{"session": {"orchestrator": "test", "context": "test"}}
).coordinator


def loader_for(coordinator, provider_mount):
async def load(module_id, config=None, **kwargs):
if module_id.startswith("provider-"):
return provider_mount

async def empty_mount(c):
return None

return empty_mount

loader = MagicMock(spec=ModuleLoader)
loader.load = AsyncMock(side_effect=load)
loader.get_on_session_ready_queue.return_value = []
coordinator.loader = loader
return loader


@pytest.mark.asyncio
async def test_partial_mount_restored_before_instance_policy_and_event(coordinator):
original = object()
await coordinator.mount("providers", original, name="openai")
observed = []

async def policy(c, spec, error):
assert c.get("providers") == {"openai": original}
assert isinstance(error, ValueError)
observed.append(copy.deepcopy(spec))
# Mutating policy input must not mutate the caller's config.
spec["config"]["default_model"] = "changed by policy"
await c.mount("providers", "unavailable", name=spec["instance_id"])

coordinator.register_capability("provider.load_failure", policy)

async def mount(c):
await c.mount("providers", object(), name="openai")
await c.mount("providers", object(), name="leaked")
raise ValueError("synthetic mount failed")

mount.__on_session_ready__ = ("provider-openai", AsyncMock())
loader = loader_for(coordinator, mount)
events = []

async def capture(event, data):
events.append(dict(data))
assert coordinator.get("providers") == {"openai": original}

coordinator.hooks.register("module:load_failed", capture, 0, name="capture")
config = {
"providers": [
{
"module": "provider-openai",
"instance_id": "broken-account",
"source": "synthetic",
"config": {"default_model": "kept", "priority": 7},
}
]
}
before = copy.deepcopy(config)
await initialize_session(config, coordinator, "test", None)
assert config == before
assert coordinator.get("providers") == {
"openai": original,
"broken-account": "unavailable",
}
assert events[0]["instance_id"] == "broken-account"
assert observed == config["providers"]
loader.enqueue_on_session_ready.assert_not_called()


@pytest.mark.asyncio
async def test_host_failure_policy_can_abort_without_observability_swallowing_it(
coordinator,
):
async def reject(*args):
raise LookupError("host chose to stop")

coordinator.register_capability("provider.load_failure", reject)

async def fail(c):
raise ValueError("failed mount")

loader_for(coordinator, fail)
with pytest.raises(LookupError, match="host chose to stop"):
await initialize_session(
{"providers": [{"module": "provider-broken"}]}, coordinator, "test", None
)


@pytest.mark.asyncio
async def test_no_policy_restores_partial_mount_and_continues_to_healthy(coordinator):
original = object()
healthy = object()
await coordinator.mount("providers", original, name="openai")
calls = 0

async def mount(c):
nonlocal calls
calls += 1
await c.mount("providers", healthy, name="openai")
if calls == 1:
raise RuntimeError("first account failed")

loader_for(coordinator, mount)
await initialize_session(
{
"providers": [
{"module": "provider-openai", "instance_id": "broken"},
{"module": "provider-openai", "instance_id": "healthy"},
]
},
coordinator,
"test",
None,
)
assert calls == 2
assert coordinator.get("providers") == {"openai": original, "healthy": healthy}


@pytest.mark.asyncio
async def test_failed_attempt_removing_default_preserves_original_order(coordinator):
first, second = object(), object()
await coordinator.mount("providers", first, name="first")
await coordinator.mount("providers", second, name="second")
before = list(coordinator.get("providers"))

async def mount(c):
await c.unmount("providers", name=before[0])
raise ValueError("removed a sibling before failing")

loader_for(coordinator, mount)
await initialize_session(
{"providers": [{"module": "provider-broken"}]}, coordinator, "test", None
)
assert list(coordinator.get("providers")) == before
assert coordinator.get("providers") == {"first": first, "second": second}


@pytest.mark.asyncio
async def test_import_failure_reports_exact_entry_and_skips_readiness(coordinator):
policy = AsyncMock()
coordinator.register_capability("provider.load_failure", policy)
loader = loader_for(coordinator, AsyncMock())
original_load = loader.load.side_effect
error = ImportError("synthetic import failure")

async def fail_load(module_id, *args, **kwargs):
if module_id == "provider-missing":
raise error
return await original_load(module_id, *args, **kwargs)

loader.load.side_effect = fail_load
spec = {
"module": "provider-missing",
"instance_id": "account",
"config": {"priority": 1},
}
await initialize_session({"providers": [spec]}, coordinator, "test", None)
policy.assert_awaited_once_with(coordinator, spec, error)
loader.enqueue_on_session_ready.assert_not_called()


@pytest.mark.asyncio
async def test_remap_failure_restores_default_and_existing_named_slot():
c = MockCoordinator()
original, named = object(), object()
await c.mount("providers", original, name="openai")
await c.mount("providers", named, name="account")
real_mount = c.mount
failures = []

async def broken_mount(point, instance, name=None):
if name == "account" and instance is not named:
# Even a remap that mutates and then fails must be restored.
await real_mount(point, instance, name=name)
raise ValueError("remap failure")
await real_mount(point, instance, name=name)

c.mount = broken_mount

async def mount(coord):
await coord.mount("providers", object(), name="openai")

mount.__on_session_ready__ = ("provider-openai", AsyncMock())
loader = loader_for(c, mount)

async def policy(coord, spec, error):
failures.append(spec["instance_id"])
assert coord.get("providers") == {"openai": original, "account": named}

c.register_capability("provider.load_failure", policy)
await initialize_session(
{"providers": [{"module": "provider-openai", "instance_id": "account"}]},
c,
"test",
None,
)
assert failures == ["account"]
loader.enqueue_on_session_ready.assert_not_called()
1 change: 1 addition & 0 deletions tests/test_session_init_module_load_failed.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ async def mount_fn(coordinator):

mock_coordinator = MagicMock()
mock_coordinator.loader = mock_loader
mock_coordinator.get_capability.return_value = None
mock_coordinator.register_cleanup = MagicMock()
mock_coordinator.get = MagicMock(return_value={})
mock_coordinator.hooks = MagicMock()
Expand Down
Loading