diff --git a/astrbot/core/provider/__init__.py b/astrbot/core/provider/__init__.py index 812e021715..0c288e4eea 100644 --- a/astrbot/core/provider/__init__.py +++ b/astrbot/core/provider/__init__.py @@ -1,4 +1,9 @@ from .entities import ProviderMetaData -from .provider import Provider, STTProvider +from .provider import Provider, STTProvider, reorder_tailing_tool_call_user -__all__ = ["Provider", "ProviderMetaData", "STTProvider"] +__all__ = [ + "Provider", + "ProviderMetaData", + "STTProvider", + "reorder_tailing_tool_call_user", +] diff --git a/astrbot/core/provider/provider.py b/astrbot/core/provider/provider.py index 891bfdea9e..21ae2137b3 100644 --- a/astrbot/core/provider/provider.py +++ b/astrbot/core/provider/provider.py @@ -2,7 +2,7 @@ import asyncio import os from collections.abc import AsyncGenerator -from typing import Literal, TypeAlias, Union +from typing import Any, Literal, TypeAlias, Union from astrbot.core.agent.message import ContentPart, Message, is_checkpoint_message from astrbot.core.agent.tool import ToolSet @@ -211,6 +211,52 @@ async def test(self, timeout: float = 45.0) -> None: ) +def _is_valid_tool_pair(asst_msg: dict[str, Any], tool_msg: dict[str, Any]) -> bool: + """判断 assistant(tool_calls) 与 tool 是否为 tool_call_id 匹配的一对。""" + if asst_msg.get("role") != "assistant" or tool_msg.get("role") != "tool": + return False + tool_calls = asst_msg.get("tool_calls") + if not tool_calls: + return False + tc_ids = {tc.get("id") for tc in tool_calls if isinstance(tc, dict)} + return tool_msg.get("tool_call_id") in tc_ids + + +def reorder_tailing_tool_call_user(messages: list[dict[str, Any]]) -> None: + """重排因伪造工具调用导致尾部 assistant(tc) → tool → user 乱序的消息。 + + 此为轻量妥协修复。未来若实现专用的上下文操作钩子或伪造工具调用钩子, + 可考虑移除此函数。 + """ + if not isinstance(messages, list) or len(messages) < 2: + return + + last = messages[-1] + if last.get("role") != "user": + return + + # 从最后一个非 user 元素向前扫描 assistant(tool_calls) + tool 成对消息 + pairs: list[tuple[dict[str, Any], dict[str, Any]]] = [] + i = len(messages) - 2 + while i >= 1: + if not _is_valid_tool_pair(messages[i - 1], messages[i]): + break + pairs.append((messages[i - 1], messages[i])) + i -= 2 + + if not pairs: + return + + # 重排:user → assistant_1, tool_1 → ... → assistant_N, tool_N + # pairs 从内到外收集(N, N-1, ..., 1),反转后按 1..N 顺序回插 + new_tail: list[dict[str, Any]] = [last] + for asst_msg, tool_msg in reversed(pairs): + new_tail.append(asst_msg) + new_tail.append(tool_msg) + + messages[:] = messages[: i + 1] + new_tail + + class STTProvider(AbstractProvider): def __init__(self, provider_config: dict, provider_settings: dict) -> None: super().__init__(provider_config) diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 27cc459622..c4cd00388b 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -14,6 +14,7 @@ from astrbot.api.provider import Provider from astrbot.core.agent.message import AudioURLPart, ContentPart, ImageURLPart, TextPart from astrbot.core.exceptions import EmptyModelOutputError +from astrbot.core.provider import reorder_tailing_tool_call_user from astrbot.core.provider.entities import LLMResponse, TokenUsage from astrbot.core.provider.func_tool_manager import ToolSet from astrbot.core.utils.media_utils import ( @@ -789,6 +790,8 @@ async def text_chat( for tool_call_result in tool_calls_result: context_query.extend(tool_call_result.to_openai_messages()) + # 伪造工具调用对需前置到用户消息之后,与真实工具调用时序对齐 + reorder_tailing_tool_call_user(context_query) system_prompt, new_messages = self._prepare_payload(context_query) model = model or self.get_model() @@ -861,6 +864,8 @@ async def text_chat_stream( for tool_call_result in tool_calls_result: context_query.extend(tool_call_result.to_openai_messages()) + # 伪造工具调用对需前置到用户消息之后,与真实工具调用时序对齐 + reorder_tailing_tool_call_user(context_query) system_prompt, new_messages = self._prepare_payload(context_query) model = model or self.get_model() diff --git a/astrbot/core/provider/sources/gemini_source.py b/astrbot/core/provider/sources/gemini_source.py index abf7bb7cf8..05dbc49ca6 100644 --- a/astrbot/core/provider/sources/gemini_source.py +++ b/astrbot/core/provider/sources/gemini_source.py @@ -17,6 +17,7 @@ from astrbot.core.agent.message import AudioURLPart, ContentPart, ImageURLPart, TextPart from astrbot.core.exceptions import EmptyModelOutputError from astrbot.core.message.message_event_result import MessageChain +from astrbot.core.provider import reorder_tailing_tool_call_user from astrbot.core.provider.entities import LLMResponse, TokenUsage from astrbot.core.provider.func_tool_manager import ToolSet from astrbot.core.utils.media_utils import ( @@ -857,6 +858,10 @@ async def text_chat( for tcr in tool_calls_result: context_query.extend(tcr.to_openai_messages()) + # 伪造工具调用对需前置到用户消息之后,与真实工具调用时序对齐。 + # Gemini API 要求 functionCall 回合必须紧跟 user 回合,否则返回 400。 + reorder_tailing_tool_call_user(context_query) + model = model or self.get_model() payloads = {"messages": context_query, "model": model} @@ -924,6 +929,10 @@ async def text_chat_stream( for tcr in tool_calls_result: context_query.extend(tcr.to_openai_messages()) + # 伪造工具调用对需前置到用户消息之后,与真实工具调用时序对齐。 + # Gemini API 要求 functionCall 回合必须紧跟 user 回合,否则返回 400。 + reorder_tailing_tool_call_user(context_query) + model = model or self.get_model() payloads = {"messages": context_query, "model": model} diff --git a/astrbot/core/provider/sources/openai_source.py b/astrbot/core/provider/sources/openai_source.py index f7870b7137..42d253d628 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -28,6 +28,7 @@ from astrbot.core.agent.tool import ToolSet from astrbot.core.exceptions import EmptyModelOutputError from astrbot.core.message.message_event_result import MessageChain +from astrbot.core.provider import reorder_tailing_tool_call_user from astrbot.core.provider.entities import LLMResponse, TokenUsage, ToolCallsResult from astrbot.core.utils.media_utils import ( describe_media_ref, @@ -528,6 +529,8 @@ def _is_empty(content: Any) -> bool: ) payloads["messages"] = final + reorder_tailing_tool_call_user(payloads["messages"]) + async def _query( self, payloads: dict, diff --git a/tests/fixtures/fake_tool_call.py b/tests/fixtures/fake_tool_call.py new file mode 100644 index 0000000000..ed996a9dcf --- /dev/null +++ b/tests/fixtures/fake_tool_call.py @@ -0,0 +1,23 @@ +"""伪造工具调用(fake tool call)共享测试数据。 + +各 provider 测试(OpenAI / Anthropic / Gemini 格式)共用同一组伪造 +assistant(tool_calls) + tool 消息对,避免场景漂移。 +""" + +FAKE_TOOL_CALL_CONTEXTS = [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "fake_recall_abc", + "type": "function", + "function": { + "name": "recall_long_term_memory", + "arguments": '{"query": "我的名字是?", "k": 5}', + }, + } + ], + }, + {"role": "tool", "tool_call_id": "fake_recall_abc", "content": "memory json"}, +] diff --git a/tests/test_anthropic_kimi_code_provider.py b/tests/test_anthropic_kimi_code_provider.py index 0dc33f58ba..b8bf055aaf 100644 --- a/tests/test_anthropic_kimi_code_provider.py +++ b/tests/test_anthropic_kimi_code_provider.py @@ -9,6 +9,7 @@ import astrbot.core.provider.sources.request_retry as request_retry from astrbot.core.exceptions import EmptyModelOutputError from astrbot.core.provider.entities import LLMResponse +from tests.fixtures.fake_tool_call import FAKE_TOOL_CALL_CONTEXTS class _FakeAsyncAnthropic: @@ -864,3 +865,75 @@ async def test_tool_choice_empty_tool_list_skips_tool_choice(monkeypatch): kwargs = _capture_payloads_create.last_kwargs assert "tools" not in kwargs assert "tool_choice" not in kwargs + + +# ── fake tool call 重排 ─────────────────────────────────────────────────────── + +_EXPECTED_REORDERED_MESSAGES = [ + {"role": "user", "content": "我的名字是?"}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "name": "recall_long_term_memory", + "input": {"query": "我的名字是?", "k": 5}, + "id": "fake_recall_abc", + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "fake_recall_abc", + "content": "memory json", + } + ], + }, +] + + +@pytest.mark.asyncio +async def test_text_chat_reorders_fake_tool_call_pair(monkeypatch): + """伪造工具调用对应重排到用户消息之后,与真实工具调用时序对齐。""" + provider = _setup_provider_with_mock_client(monkeypatch) + + await provider.text_chat(prompt="我的名字是?", contexts=FAKE_TOOL_CALL_CONTEXTS) + + assert _capture_payloads_create.last_kwargs["messages"] == ( + _EXPECTED_REORDERED_MESSAGES + ) + + +@pytest.mark.asyncio +async def test_text_chat_stream_reorders_fake_tool_call_pair(monkeypatch): + """流式路径同样应将伪造工具调用对重排到用户消息之后。""" + monkeypatch.setattr(anthropic_source, "AsyncAnthropic", _FakeAsyncAnthropic) + + provider = anthropic_source.ProviderAnthropic( + provider_config={ + "id": "anthropic-test", + "type": "anthropic_chat_completion", + "model": "claude-test", + "key": ["test-key"], + }, + provider_settings={}, + ) + + captured: dict[str, object] = {} + + async def fake_query_stream(payloads, tools, *, request_max_retries=None): + captured["messages"] = payloads["messages"] + return + yield # pragma: no cover # 保持 async generator 形态 + + monkeypatch.setattr(provider, "_query_stream", fake_query_stream) + + async for _ in provider.text_chat_stream( + prompt="我的名字是?", contexts=FAKE_TOOL_CALL_CONTEXTS + ): + pass + + assert captured["messages"] == _EXPECTED_REORDERED_MESSAGES diff --git a/tests/test_gemini_source.py b/tests/test_gemini_source.py index 9294ea46b2..e01b083d80 100644 --- a/tests/test_gemini_source.py +++ b/tests/test_gemini_source.py @@ -3,10 +3,11 @@ import httpx import pytest -from astrbot.core.exceptions import EmptyModelOutputError import astrbot.core.provider.sources.request_retry as request_retry +from astrbot.core.exceptions import EmptyModelOutputError from astrbot.core.provider.entities import LLMResponse from astrbot.core.provider.sources.gemini_source import ProviderGoogleGenAI +from tests.fixtures.fake_tool_call import FAKE_TOOL_CALL_CONTEXTS def test_gemini_empty_output_raises_empty_model_output_error(): @@ -33,6 +34,71 @@ def test_gemini_reasoning_only_output_is_allowed(): ) +def _make_provider() -> ProviderGoogleGenAI: + provider = ProviderGoogleGenAI.__new__(ProviderGoogleGenAI) + provider.api_keys = [""] + provider.model_name = "" + return provider + + +def _assert_reordered_tail(messages: list[dict]) -> None: + assert messages[-3]["role"] == "user" + assert messages[-3]["content"] == "我的名字是?" + assert messages[-2]["role"] == "assistant" + assert messages[-2]["tool_calls"][0]["id"] == "fake_recall_abc" + assert ( + messages[-2]["tool_calls"][0]["function"]["name"] == "recall_long_term_memory" + ) + assert messages[-1]["role"] == "tool" + assert messages[-1]["tool_call_id"] == "fake_recall_abc" + + +@pytest.mark.asyncio +async def test_text_chat_reorders_fake_tool_call_pair(monkeypatch): + captured = {} + + async def fake_query(payloads, tools, *, request_max_retries=None): + captured["payloads"] = payloads + return LLMResponse(role="assistant") + + provider = _make_provider() + monkeypatch.setattr(provider, "_query", fake_query) + + await provider.text_chat(prompt="我的名字是?", contexts=FAKE_TOOL_CALL_CONTEXTS) + + messages = captured["payloads"]["messages"] + _assert_reordered_tail(messages) + + # 转换后的 contents 应为 user → model(functionCall) → user(functionResponse) + contents = provider._prepare_conversation(captured["payloads"]) + assert [content.role for content in contents] == ["user", "model", "user"] + assert contents[0].parts[0].text == "我的名字是?" + assert contents[1].parts[0].function_call.name == "recall_long_term_memory" + assert contents[2].parts[0].function_response.name == "fake_recall_abc" + + +@pytest.mark.asyncio +async def test_text_chat_stream_reorders_fake_tool_call_pair(monkeypatch): + captured = {} + + async def fake_query_stream(payloads, tools, *, request_max_retries=None): + captured["payloads"] = payloads + return + yield # pragma: no cover + + provider = _make_provider() + monkeypatch.setattr(provider, "_query_stream", fake_query_stream) + + async for _ in provider.text_chat_stream( + prompt="我的名字是?", + contexts=FAKE_TOOL_CALL_CONTEXTS, + ): + pass + + messages = captured["payloads"]["messages"] + _assert_reordered_tail(messages) + + @pytest.mark.asyncio async def test_gemini_get_models_retries_transient_request_error(monkeypatch): monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0) diff --git a/tests/test_openai_source.py b/tests/test_openai_source.py index cf15e28846..7d29a784e0 100644 --- a/tests/test_openai_source.py +++ b/tests/test_openai_source.py @@ -12,9 +12,16 @@ import astrbot.core.provider.sources.openai_source as openai_source_module import astrbot.core.provider.sources.request_retry as request_retry from astrbot.core.exceptions import EmptyModelOutputError +from astrbot.core.provider import reorder_tailing_tool_call_user from astrbot.core.provider.entities import LLMResponse from astrbot.core.provider.sources.groq_source import ProviderGroq +from astrbot.core.provider.sources.longcat_source import ProviderLongCat +from astrbot.core.provider.sources.oai_aihubmix_source import ProviderAIHubMix from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial +from astrbot.core.provider.sources.openrouter_source import ProviderOpenRouter +from astrbot.core.provider.sources.xai_source import ProviderXAI +from astrbot.core.provider.sources.xiaomi_source import ProviderXiaomi +from astrbot.core.provider.sources.zhipu_source import ProviderZhipu from astrbot.core.utils.media_utils import ResolvedMediaData, file_uri_to_path @@ -2193,3 +2200,337 @@ async def fake_create(**kwargs): assert messages[1] == {"role": "user", "content": "again"} finally: await provider.terminate() + + +# ── reorder_tailing_tool_call_user ─────────────────────────────────────────── + + +def test_reorder_single_fake_pair(): + """Single fake tool call pair at tail: assistant(tc) → tool → user.""" + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "mem result"}, + {"role": "user", "content": "帮我处理"}, + ] + } + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == [ + {"role": "user", "content": "hello"}, + {"role": "user", "content": "帮我处理"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "mem result"}, + ] + + +def test_reorder_multiple_fake_pairs(): + """Multiple fake pairs from different plugins: asst₁→tool₁→asst₂→tool₂→user.""" + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "plugin_a", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "result_a"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_02", + "type": "function", + "function": {"name": "plugin_b", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_02", "content": "result_b"}, + {"role": "user", "content": "帮我处理"}, + ] + } + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == [ + {"role": "user", "content": "hello"}, + {"role": "user", "content": "帮我处理"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "plugin_a", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "result_a"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_02", + "type": "function", + "function": {"name": "plugin_b", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_02", "content": "result_b"}, + ] + + +def test_reorder_noop_when_no_fake_pair(): + """No fake pair: normal user message at tail, should be unchanged.""" + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": "next"}, + ] + } + expected = payloads["messages"][:] + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == expected + + +def test_reorder_noop_when_tool_call_id_mismatch(): + """tool_call_id does not match assistant's tool_calls — stop collecting.""" + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_99", "content": "mismatched id"}, + {"role": "user", "content": "hello"}, + ] + } + expected = payloads["messages"][:] + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == expected + + +def test_reorder_noop_when_assistant_has_no_tool_calls(): + """Assistant has no tool_calls — stop collecting.""" + payloads = { + "messages": [ + {"role": "assistant", "content": "some reply"}, + {"role": "user", "content": "hello"}, + ] + } + expected = payloads["messages"][:] + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == expected + + +def test_reorder_noop_when_no_trailing_user(): + """Tail message is not user — no-op.""" + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "reply"}, + ] + } + expected = payloads["messages"][:] + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == expected + + +def test_reorder_noop_on_empty_or_short_list(): + """Empty or too-short messages list — no-op, no crash.""" + for msgs in [None, [], [{"role": "user", "content": "hi"}]]: + reorder_tailing_tool_call_user(msgs) + + +def test_reorder_real_tool_call_not_affected(): + """Real tool call: tool → assistant(content) before user — should not reorder.""" + payloads = { + "messages": [ + {"role": "user", "content": "search plz"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "results"}, + {"role": "assistant", "content": "here are the results"}, + {"role": "user", "content": "thanks"}, + ] + } + expected = payloads["messages"][:] + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == expected + + +@pytest.mark.parametrize( + "provider_cls", + [ + ProviderGroq, + ProviderLongCat, + ProviderAIHubMix, + ProviderOpenRouter, + ProviderXAI, + ProviderXiaomi, + ProviderZhipu, + ], +) +def test_reorder_applied_through_inherited_sanitize(provider_cls): + """All OpenAI-compatible subclasses inherit _sanitize_assistant_messages + which includes the reorder fix — verify it works on each.""" + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "mem result"}, + {"role": "user", "content": "帮我处理"}, + ] + } + + provider_cls._sanitize_assistant_messages(payloads) + + assert payloads["messages"] == [ + {"role": "user", "content": "hello"}, + {"role": "user", "content": "帮我处理"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "mem result"}, + ], f"{provider_cls.__name__} did not reorder fake tool call messages" + + +def test_reorder_stops_at_first_non_pair(): + """If a non-pair message sits between pairs, only collect from the outermost contiguous block.""" + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "plugin_a", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "result_a"}, + {"role": "assistant", "content": "intermediate reply"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_02", + "type": "function", + "function": {"name": "plugin_b", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_02", "content": "result_b"}, + {"role": "user", "content": "thanks"}, + ] + } + + reorder_tailing_tool_call_user(payloads["messages"]) + + assert payloads["messages"] == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_01", + "type": "function", + "function": {"name": "plugin_a", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_01", "content": "result_a"}, + {"role": "assistant", "content": "intermediate reply"}, + {"role": "user", "content": "thanks"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_02", + "type": "function", + "function": {"name": "plugin_b", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_02", "content": "result_b"}, + ]