diff --git a/astrbot/core/astr_agent_tool_exec.py b/astrbot/core/astr_agent_tool_exec.py index 2e5915bad0..13a4c92dc9 100644 --- a/astrbot/core/astr_agent_tool_exec.py +++ b/astrbot/core/astr_agent_tool_exec.py @@ -28,6 +28,7 @@ MessageEventResult, ) from astrbot.core.platform.message_session import MessageSession +from astrbot.core.provider import Provider from astrbot.core.provider.entites import ProviderRequest from astrbot.core.provider.register import llm_tools from astrbot.core.tools.computer_tools import ( @@ -354,6 +355,36 @@ async def _execute_handoff( continue prov_settings: dict = ctx.get_config(umo=umo).get("provider_settings", {}) + fallback_ids = prov_settings.get("fallback_chat_models", []) + fallback_providers: list[Provider] = [] + if not isinstance(fallback_ids, list): + logger.warning( + "fallback_chat_models setting is not a list, skip handoff fallback providers." + ) + else: + seen_provider_ids = {prov_id} if prov_id else set() + for fallback_id in fallback_ids: + if not isinstance(fallback_id, str) or not fallback_id: + continue + if fallback_id in seen_provider_ids: + continue + fallback_provider = ctx.get_provider_by_id(fallback_id) + if fallback_provider is None: + logger.warning( + "Handoff fallback chat provider `%s` not found, skip.", + fallback_id, + ) + continue + if not isinstance(fallback_provider, Provider): + logger.warning( + "Handoff fallback chat provider `%s` is invalid type: %s, skip.", + fallback_id, + type(fallback_provider), + ) + continue + fallback_providers.append(fallback_provider) + seen_provider_ids.add(fallback_id) + agent_max_step = int(prov_settings.get("max_agent_step", 30)) stream = prov_settings.get("streaming_response", False) llm_resp = await ctx.tool_loop_agent( @@ -367,6 +398,8 @@ async def _execute_handoff( max_steps=agent_max_step, tool_call_timeout=run_context.tool_call_timeout, stream=stream, + fallback_providers=fallback_providers, + request_max_retries=prov_settings.get("request_max_retries", 5), ) yield mcp.types.CallToolResult( content=[mcp.types.TextContent(type="text", text=llm_resp.completion_text)] diff --git a/tests/unit/test_astr_agent_tool_exec.py b/tests/unit/test_astr_agent_tool_exec.py index c0a18374a5..91f643babb 100644 --- a/tests/unit/test_astr_agent_tool_exec.py +++ b/tests/unit/test_astr_agent_tool_exec.py @@ -365,6 +365,76 @@ async def _fake_tool_loop_agent(**kwargs): assert captured["tool_call_timeout"] == 120 +@pytest.mark.asyncio +async def test_execute_handoff_passes_valid_fallbacks_and_retries_to_tool_loop_agent( + monkeypatch: pytest.MonkeyPatch, +): + captured: dict = {} + + class _ChatProvider: + pass + + primary = _ChatProvider() + fallback = _ChatProvider() + non_chat_provider = object() + providers = { + "primary": primary, + "fallback": fallback, + "non-chat": non_chat_provider, + } + + async def _fake_get_current_chat_provider_id(_umo): + return "primary" + + async def _fake_tool_loop_agent(**kwargs): + captured.update(kwargs) + return SimpleNamespace(completion_text="ok") + + context = SimpleNamespace( + get_current_chat_provider_id=_fake_get_current_chat_provider_id, + get_provider_by_id=lambda provider_id: providers.get(provider_id), + tool_loop_agent=_fake_tool_loop_agent, + get_config=lambda **_kwargs: { + "provider_settings": { + "fallback_chat_models": [ + "primary", + "fallback", + "fallback", + "missing", + "non-chat", + ], + "request_max_retries": 7, + } + }, + ) + event = _DummyEvent([]) + run_context = ContextWrapper(context=SimpleNamespace(event=event, context=context)) + tool = SimpleNamespace( + name="transfer_to_subagent", + provider_id="primary", + agent=SimpleNamespace( + name="subagent", + tools=[], + instructions="subagent-instructions", + begin_dialogs=[], + run_hooks=None, + ), + ) + monkeypatch.setattr("astrbot.core.astr_agent_tool_exec.Provider", _ChatProvider) + + async for _result in FunctionToolExecutor._execute_handoff( + tool, + run_context, + image_urls_prepared=True, + input="hello", + image_urls=[], + ): + pass + + assert captured["fallback_providers"] == [fallback] + assert captured["request_max_retries"] == 7 + + @pytest.mark.asyncio async def test_background_wakeup_passes_provider_settings_to_main_agent( monkeypatch: pytest.MonkeyPatch,