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
31 changes: 28 additions & 3 deletions amplifier_foundation/spawn_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@

import asyncio
import fnmatch
import inspect
import logging
import os
import time
Expand Down Expand Up @@ -621,7 +622,7 @@ def _spec_for_instance(
for spec in provider_specs:
if not isinstance(spec, dict):
continue
spec_id = spec.get("id") or spec.get("module", "")
spec_id = spec.get("instance_id") or spec.get("id") or spec.get("module", "")
if spec_id == instance_id:
return spec
return None
Expand Down Expand Up @@ -841,7 +842,7 @@ def _resolve_provider_index(
"""
# 1. An explicit instance id is the most specific address there is.
for i, p in enumerate(providers):
if p.get("id", "") == provider_id:
if (p.get("instance_id") or p.get("id", "")) == provider_id:
return i

# 2. Otherwise the name addresses a MODULE -- gather every instance of it.
Expand Down Expand Up @@ -952,7 +953,7 @@ def _build_provider_lookup(
# instance addressable by its own id even if its id never appears as a
# module-type key above.)
for i, p in enumerate(providers):
instance_id = p.get("id")
instance_id = p.get("instance_id") or p.get("id")
if instance_id:
lookup[instance_id] = i

Expand Down Expand Up @@ -1141,6 +1142,10 @@ async def apply_provider_preferences_with_resolution(

providers = mount_plan.get("providers", [])
attempt_results = diagnostics if diagnostics is not None else []
getter = getattr(coordinator, "get_capability", None)
availability = getter("provider.check_available") if callable(getter) else None
availability = availability if callable(availability) else None
failures = []

# Find first matching preference whose model actually resolves, and
# apply it. A preference whose provider is present but whose glob
Expand All @@ -1166,6 +1171,20 @@ async def apply_provider_preferences_with_resolution(

target = providers[target_idx]
runtime_provider = _runtime_provider_name(target)
if availability is not None:
try:
checked = availability(runtime_provider)
except Exception as error:
failures.append(error)
attempt_results.append(ModelResolutionResult(
resolved_model=None, pattern=pref.model,
status="provider_unavailable", provider=runtime_provider,
))
continue # Only another declared preference may supply fallback.
if inspect.isawaitable(checked):
if inspect.iscoroutine(checked):
checked.close()
raise TypeError("provider.check_available must be synchronous; no provider was selected.")
# Resolve model pattern if it's a glob. The runtime instance comes
# from the same model-aware selection that chose target_idx: never
# query a bare alias then apply the resulting model to another mount.
Expand Down Expand Up @@ -1196,6 +1215,12 @@ async def apply_provider_preferences_with_resolution(
# in the mount plan, or every candidate's model pattern failed to
# resolve. Either way, leave the mount plan unmodified rather than
# writing an unresolved pattern string into it.
if availability is not None:
if failures:
raise failures[0]
# Under the host's isolation policy an explicit preference cannot
# quietly turn into the parent/default account when resolution fails.
raise ValueError("No declared provider preference could be resolved; the default provider was not substituted.")
if diagnostics is None:
statuses = ", ".join(sorted({result.status for result in attempt_results}))
logger.warning(
Expand Down
16 changes: 16 additions & 0 deletions docs/provider-failure-policy.md
Original file line number Diff line number Diff line change
Expand Up @@ -59,3 +59,19 @@ 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.

## Applying a selected child's provider preferences

A host that retains unavailable accounts may additionally register the
synchronous `provider.check_available(instance_id)` capability. Returning means
the mounted account is available; raising preserves its actionable failure.
`apply_provider_preferences_with_resolution` calls it for the exact configured
account before an exact model or catalog lookup, honoring Core `instance_id`
before the older `id` field. Only another entry in the declared preference list
may provide fallback. If all entries fail or no model resolves, it raises rather
than leaving the default account selected. Cancellation is never fallback.
Without this capability, legacy no-match behavior remains unchanged.

This is a host boundary, not an automatic health probe. The host owns safe
error text, per-session account state, revalidation, and applying this policy on
each root/resumed/child session. Never publish provider secrets in exceptions.
3 changes: 3 additions & 0 deletions tests/test_named_delegate_matrix_67u.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,9 @@ def _make_coordinator(
p["id"]: _FakeProvider(models_by_module.get(p["module"], [])) for p in providers
}
coordinator = MagicMock()
# These characterize the legacy host, which has not opted into an
# availability policy. An unconfigured MagicMock invents a callable.
coordinator.get_capability.return_value = None
coordinator.config = {"providers": providers}
coordinator.get = MagicMock(
side_effect=lambda key: runtime if key == "providers" else None
Expand Down
99 changes: 99 additions & 0 deletions tests/test_provider_availability.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
"""Explicit child preferences must not revert to a different default account."""
import copy
from types import SimpleNamespace
from unittest.mock import AsyncMock

import pytest

from amplifier_foundation.spawn_utils import (
ProviderPreference,
_build_provider_lookup,
_find_provider_instance,
apply_provider_preferences_with_resolution,
)


def fixture_host(broken=()):
specs = [
{"module": "provider-openai", "instance_id": "work", "config": {"default_model": "work-model", "priority": 1}},
{"module": "provider-openai", "instance_id": "personal", "config": {"default_model": "personal-model", "priority": 0}},
]
providers = {name: SimpleNamespace(list_models=AsyncMock(return_value=[name + "-model"]))
for name in ("work", "personal")}
checked = []

def check(name):
checked.append(name)
if name in broken:
raise ValueError("Unavailable configured account: " + name)

coordinator = SimpleNamespace(config={"providers": specs}, get=lambda _: providers,
get_capability=lambda key: check if key == "provider.check_available" else None)
return specs, providers, coordinator, checked


@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["work-model", "work-*"])
async def test_failed_explicit_account_never_uses_healthy_default(model):
specs, providers, coordinator, checked = fixture_host({"work"})
plan = {"providers": specs}
original = copy.deepcopy(plan)
with pytest.raises(ValueError, match="account: work"):
await apply_provider_preferences_with_resolution(plan, [ProviderPreference("work", model)], coordinator)
assert plan == original and checked == ["work"]
assert _build_provider_lookup(specs)["work"] == 0
for provider in providers.values():
provider.list_models.assert_not_awaited()


@pytest.mark.asyncio
async def test_declared_preference_fallback_keeps_exact_account_and_config():
specs, providers, coordinator, checked = fixture_host({"work"})
original = copy.deepcopy(specs)
prefs = [ProviderPreference("work", "work-*"), ProviderPreference("personal", "personal-*")]
result = await apply_provider_preferences_with_resolution({"providers": specs}, prefs, coordinator)
assert checked == ["work", "personal"]
assert result["providers"][1]["config"]["default_model"] == "personal-model"
assert specs == original
assert _find_provider_instance(providers, "openai", coordinator) is providers["personal"]


@pytest.mark.asyncio
@pytest.mark.parametrize("provider,model", [("work", "missing-*"), ("missing", "anything")])
async def test_no_matching_preference_under_host_policy_cannot_keep_default(provider, model):
specs, _, coordinator, _ = fixture_host()
with pytest.raises(ValueError, match="default provider was not substituted"):
await apply_provider_preferences_with_resolution({"providers": specs}, [ProviderPreference(provider, model)], coordinator)


@pytest.mark.asyncio
async def test_no_policy_preserves_legacy_no_match_behavior():
specs, _, coordinator, _ = fixture_host()
coordinator.get_capability = lambda _: None
plan = {"providers": specs}
assert await apply_provider_preferences_with_resolution(plan, [ProviderPreference("missing", "anything")], coordinator) is plan


@pytest.mark.asyncio
async def test_cancellation_is_not_a_candidate_fallback():
import asyncio
specs, _, coordinator, _ = fixture_host()
def cancel(_):
raise asyncio.CancelledError()
coordinator.get_capability = lambda _: cancel
with pytest.raises(asyncio.CancelledError):
await apply_provider_preferences_with_resolution({"providers": specs},
[ProviderPreference("work", "work-model"), ProviderPreference("personal", "personal-model")], coordinator)


@pytest.mark.asyncio
async def test_async_policy_is_a_contract_error_not_permission_or_fallback():
specs, providers, coordinator, _ = fixture_host()
async def invalid(_):
raise ValueError("must never run")
coordinator.get_capability = lambda _: invalid
with pytest.raises(TypeError, match="must be synchronous"):
await apply_provider_preferences_with_resolution({"providers": specs},
[ProviderPreference("work", "work-*"), ProviderPreference("personal", "personal-*")], coordinator)
for provider in providers.values():
provider.list_models.assert_not_awaited()
2 changes: 2 additions & 0 deletions tests/test_spawn_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -691,6 +691,7 @@ async def test_all_preferences_fail_leaves_mount_plan_unmodified(self) -> None:
mock_openai.list_models = AsyncMock(return_value=["gpt-4o", "gpt-4o-mini"])

mock_coordinator = MagicMock()
mock_coordinator.get_capability.return_value = None
mock_coordinator.get.return_value = {
"provider-anthropic": mock_anthropic,
"provider-openai": mock_openai,
Expand Down Expand Up @@ -783,6 +784,7 @@ async def test_query_failure_never_uses_persisted_default_as_a_fallback(self) ->
]
}
coordinator = MagicMock()
coordinator.get_capability.return_value = None
coordinator.get.side_effect = RuntimeError("network details must stay hidden")
diagnostics = []

Expand Down
Loading