|
14 | 14 | from astrbot.core.agent.agent import Agent |
15 | 15 | from astrbot.core.agent.handoff import HandoffTool |
16 | 16 | from astrbot.core.agent.hooks import BaseAgentRunHooks |
17 | | -from astrbot.core.agent.message import ImageURLPart, Message, TextPart |
| 17 | +from astrbot.core.agent.message import ( |
| 18 | + AssistantMessageSegment, |
| 19 | + ImageURLPart, |
| 20 | + Message, |
| 21 | + TextPart, |
| 22 | + ToolCallMessageSegment, |
| 23 | +) |
18 | 24 | from astrbot.core.agent.run_context import ContextWrapper |
19 | 25 | from astrbot.core.agent.runners.tool_loop_agent_runner import ToolLoopAgentRunner |
20 | 26 | from astrbot.core.agent.tool import FunctionTool, ToolSet |
21 | 27 | from astrbot.core.astr_agent_tool_exec import FunctionToolExecutor |
22 | 28 | from astrbot.core.exceptions import EmptyModelOutputError |
23 | | -from astrbot.core.provider.entities import LLMResponse, ProviderRequest, TokenUsage |
| 29 | +from astrbot.core.provider.entities import ( |
| 30 | + LLMResponse, |
| 31 | + ProviderRequest, |
| 32 | + TokenUsage, |
| 33 | + ToolCallsResult, |
| 34 | +) |
24 | 35 | from astrbot.core.provider.provider import Provider |
25 | 36 |
|
26 | 37 |
|
@@ -1559,6 +1570,156 @@ async def test_tool_result_injects_follow_up_notice( |
1559 | 1570 | assert ticket2.consumed is True |
1560 | 1571 |
|
1561 | 1572 |
|
| 1573 | +@pytest.mark.asyncio |
| 1574 | +async def test_reset_appends_injected_tool_calls_result_after_user( |
| 1575 | + runner, mock_provider, mock_tool_executor, mock_hooks |
| 1576 | +): |
| 1577 | + """reset() 把 on_llm_request 注入的 tool_calls_result 追加在当前 user 消息之后。 |
| 1578 | +
|
| 1579 | + 最终顺序:history → 当前 user → assistant(tool_calls) → tool(result)。 |
| 1580 | + """ |
| 1581 | + request = ProviderRequest( |
| 1582 | + prompt="当前问题", |
| 1583 | + contexts=[ |
| 1584 | + {"role": "user", "content": "历史消息"}, |
| 1585 | + {"role": "assistant", "content": "历史回答"}, |
| 1586 | + ], |
| 1587 | + tool_calls_result=ToolCallsResult( |
| 1588 | + tool_calls_info=AssistantMessageSegment( |
| 1589 | + tool_calls=[ |
| 1590 | + { |
| 1591 | + "id": "fake_1", |
| 1592 | + "type": "function", |
| 1593 | + "function": { |
| 1594 | + "name": "recall_long_term_memory", |
| 1595 | + "arguments": "{}", |
| 1596 | + }, |
| 1597 | + } |
| 1598 | + ] |
| 1599 | + ), |
| 1600 | + tool_calls_result=[ |
| 1601 | + ToolCallMessageSegment( |
| 1602 | + tool_call_id="fake_1", |
| 1603 | + content="memory json", |
| 1604 | + ) |
| 1605 | + ], |
| 1606 | + ), |
| 1607 | + ) |
| 1608 | + |
| 1609 | + await runner.reset( |
| 1610 | + provider=mock_provider, |
| 1611 | + request=request, |
| 1612 | + run_context=ContextWrapper(context=None), |
| 1613 | + tool_executor=mock_tool_executor, |
| 1614 | + agent_hooks=mock_hooks, |
| 1615 | + streaming=False, |
| 1616 | + ) |
| 1617 | + |
| 1618 | + roles = [m.role for m in runner.run_context.messages] |
| 1619 | + assert roles == ["user", "assistant", "user", "assistant", "tool"] |
| 1620 | + assert runner.run_context.messages[-2].tool_calls[0].id == "fake_1" |
| 1621 | + assert runner.run_context.messages[-1].tool_call_id == "fake_1" |
| 1622 | + |
| 1623 | + |
| 1624 | +@pytest.mark.asyncio |
| 1625 | +async def test_runner_with_openai_provider_preserves_injected_tool_calls_order( |
| 1626 | + mock_tool_executor, mock_hooks |
| 1627 | +): |
| 1628 | + """端到端:runner 消费注入的 tool_calls_result 后,OpenAI payload 顺序保持正确。 |
| 1629 | +
|
| 1630 | + 覆盖 ToolLoopAgentRunner → ProviderOpenAIOfficial.text_chat → _query 的完整 |
| 1631 | + payload 组装路径(真实 provider,仅对 SDK create 的入口 _query 打桩捕获)。 |
| 1632 | + 顺序:history → 当前 user → assistant(tool_calls) → tool(result)。 |
| 1633 | + """ |
| 1634 | + from astrbot.core.agent.tool import FunctionTool, ToolSet |
| 1635 | + from astrbot.core.provider.entities import ToolCallsResult |
| 1636 | + from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial |
| 1637 | + |
| 1638 | + provider = ProviderOpenAIOfficial( |
| 1639 | + provider_config={ |
| 1640 | + "id": "test-openai", |
| 1641 | + "type": "openai_chat_completion", |
| 1642 | + "model": "gpt-4o-mini", |
| 1643 | + "key": ["test-key"], |
| 1644 | + }, |
| 1645 | + provider_settings={}, |
| 1646 | + ) |
| 1647 | + captured: dict[str, list[dict[str, Any]]] = {} |
| 1648 | + |
| 1649 | + async def fake_query(payloads, func_tool, *, request_max_retries=None): |
| 1650 | + captured["messages"] = [dict(m) for m in payloads["messages"]] |
| 1651 | + return LLMResponse(role="assistant", completion_text="ok") |
| 1652 | + |
| 1653 | + provider._query = fake_query # type: ignore[method-assign] |
| 1654 | + |
| 1655 | + tool_set = ToolSet( |
| 1656 | + tools=[ |
| 1657 | + FunctionTool( |
| 1658 | + name="test_tool", |
| 1659 | + description="测试工具", |
| 1660 | + parameters={ |
| 1661 | + "type": "object", |
| 1662 | + "properties": {"query": {"type": "string"}}, |
| 1663 | + }, |
| 1664 | + handler=AsyncMock(), |
| 1665 | + ) |
| 1666 | + ] |
| 1667 | + ) |
| 1668 | + request = ProviderRequest( |
| 1669 | + prompt="当前问题", |
| 1670 | + func_tool=tool_set, |
| 1671 | + contexts=[ |
| 1672 | + {"role": "user", "content": "历史消息"}, |
| 1673 | + {"role": "assistant", "content": "历史回答"}, |
| 1674 | + ], |
| 1675 | + tool_calls_result=ToolCallsResult( |
| 1676 | + tool_calls_info=AssistantMessageSegment( |
| 1677 | + tool_calls=[ |
| 1678 | + { |
| 1679 | + "id": "fake_1", |
| 1680 | + "type": "function", |
| 1681 | + "function": { |
| 1682 | + "name": "recall_long_term_memory", |
| 1683 | + "arguments": "{}", |
| 1684 | + }, |
| 1685 | + } |
| 1686 | + ] |
| 1687 | + ), |
| 1688 | + tool_calls_result=[ |
| 1689 | + ToolCallMessageSegment( |
| 1690 | + tool_call_id="fake_1", |
| 1691 | + content="memory json", |
| 1692 | + ) |
| 1693 | + ], |
| 1694 | + ), |
| 1695 | + ) |
| 1696 | + |
| 1697 | + runner = ToolLoopAgentRunner() |
| 1698 | + try: |
| 1699 | + await runner.reset( |
| 1700 | + provider=provider, |
| 1701 | + request=request, |
| 1702 | + run_context=ContextWrapper(context=None), |
| 1703 | + tool_executor=mock_tool_executor, |
| 1704 | + agent_hooks=mock_hooks, |
| 1705 | + streaming=False, |
| 1706 | + ) |
| 1707 | + async for _ in runner.step(): |
| 1708 | + pass |
| 1709 | + finally: |
| 1710 | + await provider.terminate() |
| 1711 | + |
| 1712 | + assert [m["role"] for m in captured["messages"]] == [ |
| 1713 | + "user", |
| 1714 | + "assistant", |
| 1715 | + "user", |
| 1716 | + "assistant", |
| 1717 | + "tool", |
| 1718 | + ] |
| 1719 | + assert captured["messages"][-2]["tool_calls"][0]["id"] == "fake_1" |
| 1720 | + assert captured["messages"][-1]["tool_call_id"] == "fake_1" |
| 1721 | + |
| 1722 | + |
1562 | 1723 | @pytest.mark.asyncio |
1563 | 1724 | async def test_follow_up_ticket_not_consumed_when_no_next_tool_call( |
1564 | 1725 | runner, mock_provider, provider_request, mock_tool_executor, mock_hooks |
|
0 commit comments