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
92 changes: 81 additions & 11 deletions amplifier_module_loop_streaming/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2756,7 +2756,14 @@ async def _judge_stall(
)

request_result = await hooks.emit(
PROVIDER_REQUEST, {"provider": provider_name, "iteration": 0}
PROVIDER_REQUEST,
{
"provider": provider_name,
"model": model_override
if model_override is not None
else self._provider_default_model(provider),
"iteration": 0,
},
)
if coordinator:
request_result = await coordinator.process_hook_result(
Expand Down Expand Up @@ -3021,7 +3028,14 @@ async def _summarize_goal_run(
)

request_result = await hooks.emit(
PROVIDER_REQUEST, {"provider": provider_name, "iteration": 0}
PROVIDER_REQUEST,
{
"provider": provider_name,
"model": model_override
if model_override is not None
else self._provider_default_model(provider),
"iteration": 0,
},
)
if coordinator:
request_result = await coordinator.process_hook_result(
Expand Down Expand Up @@ -3193,7 +3207,14 @@ async def _evaluate_goal(
# that gate LLM calls (approval, cost-aware routing, rate limiting)
# see and can govern the evaluator's call too.
request_result = await hooks.emit(
PROVIDER_REQUEST, {"provider": provider_name, "iteration": 0}
PROVIDER_REQUEST,
{
"provider": provider_name,
"model": model_override
if model_override is not None
else self._provider_default_model(provider),
"iteration": 0,
},
)
if coordinator:
request_result = await coordinator.process_hook_result(
Expand Down Expand Up @@ -3513,12 +3534,10 @@ async def request_messages(
yield ("Error: No providers available", 0)
return

# Find provider name for event emission
provider_name = None
for name, prov in providers.items():
if prov is provider:
provider_name = name
break
# Host wrappers can select a request-specific view of a mounted provider.
# Keep that view for execution, but report the exact mounted alias so
# provider-request hooks resolve the same provider before building tools.
provider_name = self._provider_name(provider, providers)
budget_capable = callable(getattr(provider, "request_budget", None))
measured_compaction_capable = (
measured_view_getter is not None
Expand Down Expand Up @@ -4222,7 +4241,12 @@ async def recover_context_overflow(
if self._reminder_placement == "pre_user":
turn_start_result = await hooks.emit(
PROVIDER_REQUEST,
{"provider": provider_name, "iteration": 1, "phase": "turn_start"},
{
"provider": provider_name,
"model": self._provider_default_model(provider),
"iteration": 1,
"phase": "turn_start",
},
)
if coordinator:
turn_start_result = await coordinator.process_hook_result(
Expand Down Expand Up @@ -4422,7 +4446,11 @@ async def exit_for_cancellation() -> None:
else:
result = await hooks.emit(
PROVIDER_REQUEST,
{"provider": provider_name, "iteration": iteration},
{
"provider": provider_name,
"model": self._provider_default_model(provider),
"iteration": iteration,
},
)
if coordinator:
result = await coordinator.process_hook_result(
Expand Down Expand Up @@ -5549,6 +5577,7 @@ async def count_view(
PROVIDER_REQUEST,
{
"provider": provider_name,
"model": self._provider_default_model(provider),
"iteration": iteration,
"max_reached": True,
},
Expand Down Expand Up @@ -6739,6 +6768,47 @@ async def _process_tools(self, context, tools, hooks) -> None:
"""Process any pending tool calls."""
# Simplified - would process tracked tool calls

@staticmethod
def _provider_name(provider: Any, providers: dict[str, Any]) -> str | None:
"""Resolve a mounted alias through transparent host wrappers by identity.

``original`` and ``__wrapped__`` retain provider provenance without
replacing the mounted objects. Never match vendor ids or display names:
several instances can serve the same vendor with different credentials
and defaults. Prefer an exact match and decline ambiguous provenance.
This is recomputed each turn so replacing a host selection stays valid.
"""
for name, mounted in providers.items():
if mounted is provider:
return name

def identities(value: Any) -> set[int]:
seen: set[int] = set()
pending = [value]
while pending and len(seen) < 32:
current = pending.pop()
if id(current) in seen:
continue
seen.add(id(current))
for attribute in ("original", "__wrapped__"):
# Read the wrapper's own link, not a delegated __getattr__
# that may skip layers or recurse on malformed cycles.
try:
wrapped = object.__getattribute__(current, attribute)
except AttributeError:
continue
if wrapped is not None and id(wrapped) not in seen:
pending.append(wrapped)
return seen

selected = identities(provider)
matches = [
name
for name, mounted in providers.items()
if selected.intersection(identities(mounted))
]
return matches[0] if len(matches) == 1 else None

def _select_provider(self, providers: dict[str, Any]) -> Any:
"""Select the provider for the TOP-LEVEL CONVERSATION.

Expand Down
123 changes: 123 additions & 0 deletions tests/test_provider_pin.py
Original file line number Diff line number Diff line change
Expand Up @@ -718,3 +718,126 @@ def test_not_mounted_check_precedes_vendor_check() -> None:

with pytest.raises(ValueError, match="not mounted"):
pin.pin("gemini-pro")


class _ProviderView:
"""Transparent request/telemetry views used by application hosts."""

def __init__(self, original):
self.original = original

def __getattr__(self, name):
return getattr(self.original, name)


class _SelectedProviderView(_ProviderView):
def __init__(self, original, model):
super().__init__(original)
self.model = model

def get_info(self):
info = self.original.get_info()
return info.model_copy(
update={"defaults": {**info.defaults, "model": self.model}}
)

async def complete(self, request, **kwargs):
return await self.original.complete(
request.model_copy(update={"model": self.model}), **kwargs
)


def test_provider_alias_uses_identity_through_independent_wrapper_chains():
first = StubProvider("Anthropic", default_model="first")
second = StubProvider("Anthropic", default_model="second")
providers = {
"first": _ProviderView(first),
"second": _ProviderView(_ProviderView(second)),
}
assert (
StreamingOrchestrator._provider_name(
_SelectedProviderView(second, "override"), providers
)
== "second"
)
# Same vendor metadata is not evidence of the same mounted instance.
assert (
StreamingOrchestrator._provider_name(StubProvider("Anthropic"), providers)
is None
)
# Cycles terminate and ambiguous aliases are never guessed.
cycle = _ProviderView(second)
cycle.original = cycle
assert StreamingOrchestrator._provider_name(cycle, providers) is None
assert (
StreamingOrchestrator._provider_name(
_ProviderView(second), {"a": second, "b": second}
)
is None
)
assert StreamingOrchestrator._provider_name(second, {"a": second}) == "a"


@pytest.mark.asyncio
@pytest.mark.parametrize("placement", ["pre_user", "tail"])
async def test_wrapped_selection_reports_alias_and_effective_model_before_tool_snapshot(
placement,
):
first = StubProvider("Anthropic", default_model="first-default")
second = StubProvider("Anthropic", default_model="second-default")
providers = {"first": _ProviderView(first), "second": _ProviderView(second)}
original_mounts = dict(providers)
selected = _SelectedProviderView(second, "second-override")
orch = StreamingOrchestrator(
{"max_iterations": 5, "stream_delay": 0, "reminder_placement": placement}
)
orch._select_provider = lambda _providers: selected
events = []
snapshots = []
hooks = HookRegistry()

async def before_request(event, data):
events.append(dict(data))

class Tool:
name = "probe"
description = "Capture native-tool construction timing"
@property
def input_schema(self):
return {"type": "object", "properties": {}}

@property
def native_tool_spec(self):
snapshots.append(dict(events[-1]))
return {"type": "test-native-tool"}

hooks.register("provider:request", before_request)
for underlying, model, alias in (
(second, "second-override", "second"),
(first, "first-override", "first"),
):
# A warmed application replaces the root selection while retaining the
# same prepared execution map. The metadata must follow each new turn.
selected = _SelectedProviderView(underlying, model)
await orch.execute(
prompt="hi",
context=StubContext(),
providers=providers,
tools={"probe": Tool()},
hooks=hooks,
)
assert events[-1]["provider"] == alias
assert events[-1]["model"] == model
assert snapshots[-1]["provider"] == alias
assert snapshots[-1]["model"] == model
assert providers == original_mounts
assert first.get_info().defaults["model"] == "first-default"
assert second.get_info().defaults["model"] == "second-default"
# Goal/worker routing sees the unchanged provider map and original defaults.
orch._goal_model_cache = None
name, utility, model, config = await orch._resolve_goal_model(
providers, StubCoordinator(providers)
)
assert name == "first" and utility is providers["first"]
assert model is None and config == {}
assert utility.get_info().defaults["model"] == "first-default"
Loading