From 79e2eb3cdbfce0bd492b9f2258feb2276ec21d3d Mon Sep 17 00:00:00 2001 From: Brian Krabach Date: Wed, 23 Sep 2026 23:39:33 -0700 Subject: [PATCH] Expose host provider preflight before loading configured accounts --- docs/CAPABILITY_REGISTRY.md | 10 +++++++ python/amplifier_core/_session_init.py | 3 ++ tests/test_provider_failure_policy.py | 41 ++++++++++++++++++++++++++ 3 files changed, 54 insertions(+) diff --git a/docs/CAPABILITY_REGISTRY.md b/docs/CAPABILITY_REGISTRY.md index 1552dd9..5839cf8 100644 --- a/docs/CAPABILITY_REGISTRY.md +++ b/docs/CAPABILITY_REGISTRY.md @@ -29,10 +29,20 @@ 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.before_load` | `async (coordinator, provider_spec: dict) → None` | host application | Python session initialization | | `provider.load_failure` | `async (coordinator, provider_spec: dict, error: Exception) → None` | host application | Python session initialization | ### Provider initialization failures +`provider.before_load` is an optional asynchronous host preflight. It receives +an independent copy of the exact configured entry before loader or mount code +runs. Returning continues normal loading; raising routes through the same mount +restoration, observability event and `provider.load_failure` callback as an +import or mount failure. It lets a host retain a provider whose configuration +failed validation without attempting it with incomplete or substituted credentials. +Neither callback changes the configured entry unless the host explicitly edits +its own session state. + 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 diff --git a/python/amplifier_core/_session_init.py b/python/amplifier_core/_session_init.py index cb77c1c..46474cf 100644 --- a/python/amplifier_core/_session_init.py +++ b/python/amplifier_core/_session_init.py @@ -167,6 +167,9 @@ async def initialize_session( instance_id = provider_config.get("instance_id") # multi-instance support before_providers = dict(coordinator.get("providers") or {}) try: + before_load = coordinator.get_capability("provider.before_load") + if before_load is not None: + await before_load(coordinator, deepcopy(provider_config)) logger.info( f"Loading provider: {module_id}" + (f" (instance: {instance_id})" if instance_id else "") diff --git a/tests/test_provider_failure_policy.py b/tests/test_provider_failure_policy.py index c18079d..a1c16ed 100644 --- a/tests/test_provider_failure_policy.py +++ b/tests/test_provider_failure_policy.py @@ -218,3 +218,44 @@ async def policy(coord, spec, error): ) assert failures == ["account"] loader.enqueue_on_session_ready.assert_not_called() + + +@pytest.mark.asyncio +async def test_host_preflight_rejects_exact_account_before_loader(coordinator): + before_calls = [] + failures = [] + + async def preflight(c, spec): + before_calls.append(spec["instance_id"]) + if spec["instance_id"] == "invalid": + spec["config"]["api_key"] = "callback-local-only" + raise ValueError("host preflight failed") + + async def failed(c, spec, error): + failures.append(spec) + assert spec["config"]["api_key"] == "synthetic-value" + + coordinator.register_capability("provider.before_load", preflight) + coordinator.register_capability("provider.load_failure", failed) + + async def mount(c): + await c.mount("providers", "healthy", name="x") + + loader = loader_for(coordinator, mount) + specs = [ + { + "module": "provider-x", + "instance_id": "invalid", + "config": {"api_key": "synthetic-value"}, + }, + {"module": "provider-x", "instance_id": "valid", "config": {}}, + ] + original = copy.deepcopy(specs) + await initialize_session({"providers": specs}, coordinator, "test", None) + assert specs == original + assert before_calls == ["invalid", "valid"] + assert failures == specs[:1] + assert coordinator.get("providers") == {"valid": "healthy"} + assert [call.args[0] for call in loader.load.await_args_list].count( + "provider-x" + ) == 1