diff --git a/docs/CAPABILITY_REGISTRY.md b/docs/CAPABILITY_REGISTRY.md index cc6a4b5..1552dd9 100644 --- a/docs/CAPABILITY_REGISTRY.md +++ b/docs/CAPABILITY_REGISTRY.md @@ -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 diff --git a/docs/module-failure-reasons.md b/docs/module-failure-reasons.md index f987a01..0fc6249 100644 --- a/docs/module-failure-reasons.md +++ b/docs/module-failure-reasons.md @@ -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. diff --git a/python/amplifier_core/_session_init.py b/python/amplifier_core/_session_init.py index 98002b3..cb77c1c 100644 --- a/python/amplifier_core/_session_init.py +++ b/python/amplifier_core/_session_init.py @@ -7,6 +7,7 @@ """ import logging +from copy import deepcopy from typing import Any logger = logging.getLogger(__name__) @@ -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 @@ -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: @@ -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}" @@ -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: @@ -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", []): diff --git a/tests/test_provider_failure_policy.py b/tests/test_provider_failure_policy.py new file mode 100644 index 0000000..c18079d --- /dev/null +++ b/tests/test_provider_failure_policy.py @@ -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() diff --git a/tests/test_session_init_module_load_failed.py b/tests/test_session_init_module_load_failed.py index ce7c12c..4578b37 100644 --- a/tests/test_session_init_module_load_failed.py +++ b/tests/test_session_init_module_load_failed.py @@ -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()