Skip to content
Open
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
9 changes: 7 additions & 2 deletions astrbot/core/provider/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
48 changes: 47 additions & 1 deletion astrbot/core/provider/provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
5 changes: 5 additions & 0 deletions astrbot/core/provider/sources/anthropic_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down
9 changes: 9 additions & 0 deletions astrbot/core/provider/sources/gemini_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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}
Expand Down Expand Up @@ -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}
Expand Down
3 changes: 3 additions & 0 deletions astrbot/core/provider/sources/openai_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
23 changes: 23 additions & 0 deletions tests/fixtures/fake_tool_call.py
Original file line number Diff line number Diff line change
@@ -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"},
]
73 changes: 73 additions & 0 deletions tests/test_anthropic_kimi_code_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
68 changes: 67 additions & 1 deletion tests/test_gemini_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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)
Expand Down
Loading
Loading