diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 0fa9b69..54c1a59 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -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( @@ -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( @@ -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( @@ -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 @@ -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( @@ -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( @@ -5549,6 +5577,7 @@ async def count_view( PROVIDER_REQUEST, { "provider": provider_name, + "model": self._provider_default_model(provider), "iteration": iteration, "max_reached": True, }, @@ -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. diff --git a/tests/test_provider_pin.py b/tests/test_provider_pin.py index f8b4e94..563384e 100644 --- a/tests/test_provider_pin.py +++ b/tests/test_provider_pin.py @@ -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"