feat(agent): stream non-verbose tool progress

This commit is contained in:
jxxghp
2026-09-13 00:22:58 +08:00
parent 0c28ed82da
commit f70d7c1410
3 changed files with 199 additions and 18 deletions

View File

@@ -42,7 +42,7 @@ class StreamingHandler:
- 后续有新内容时编辑同一条消息(通过 edit_message
- 当消息长度接近渠道限制时,冻结当前消息并发送新消息继续输出
4. 工具调用时:
- 流式渠道:工具消息直接 emit() 追加到 buffer与 Agent 文字合并为同一条流式消息
- 可编辑渠道:工具摘要实时写入 buffer同一批次的后续调用会原地更新摘要块
- 非流式渠道:调用 take() 取出已积累的文字,与工具消息合并独立发送
5. Agent最终完成时调用 stop_streaming():执行最后一次刷新,
返回是否已通过流式发送完所有内容(调用方据此决定是否还需额外发送)
@@ -75,8 +75,13 @@ class StreamingHandler:
self._original_chat_id: Optional[str] = None
self._title: str = ""
self._allow_dispatch_without_context = False
# 非啰嗦模式下的待输出工具统计,等下一段文本到来时再统一补一句摘要
# 非啰嗦模式下尚未写入展示的工具统计。
self._pending_tool_stats: dict[str, dict[str, Any]] = {}
# 当前正在实时更新的工具摘要块:(起始偏移, 完整块, 摘要行)。
# 正文到来后清空,使下一次工具调用开启新的统计批次。
self._live_tool_summary: Optional[tuple[int, str, str]] = None
# 当前实时摘要批次的累计统计,不能随着每次渲染摘要被消费掉。
self._live_tool_stats: dict[str, dict[str, Any]] = {}
# 本轮已写入缓冲区的工具展示行,供 Telegram 富文本渲染时做区分样式
self._tool_summaries: set[str] = set()
@@ -96,8 +101,16 @@ class StreamingHandler:
emitted = token or ""
if self._pending_tool_stats:
summary = self._consume_pending_tool_summary_locked()
if self._live_tool_summary:
# 实时摘要已经直接更新了 buffer正文只需继续追加避免把
# 本次替换后的摘要块再次拼接到缓冲区末尾。
self._flush_live_tool_summary_locked()
summary = ""
else:
summary = self._consume_pending_tool_summary_locked()
if summary:
self._live_tool_summary = None
self._live_tool_stats = {}
if emitted:
emitted = f"{summary}{emitted.lstrip(chr(10))}"
else:
@@ -107,6 +120,9 @@ class StreamingHandler:
if self._buffer.endswith("\n\n") and emitted.startswith("\n"):
emitted = emitted.lstrip("\n")
self._buffer += emitted
if emitted:
self._live_tool_summary = None
self._live_tool_stats = {}
return emitted
def emit_tool_message(self, message: str) -> str:
@@ -179,10 +195,14 @@ class StreamingHandler:
with self._lock:
if not self._buffer:
self._live_tool_summary = None
self._live_tool_stats = {}
return ""
message = self._buffer
logger.info(f"Agent消息: {message}")
self._buffer = ""
self._live_tool_summary = None
self._live_tool_stats = {}
return message
def clear(self):
@@ -196,6 +216,8 @@ class StreamingHandler:
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._tool_summaries = set()
self._live_tool_summary = None
self._live_tool_stats = {}
def reset(self):
"""
@@ -211,6 +233,8 @@ class StreamingHandler:
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._tool_summaries = set()
self._live_tool_summary = None
self._live_tool_stats = {}
async def start_streaming(
self,
@@ -274,6 +298,8 @@ class StreamingHandler:
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._tool_summaries = set()
self._live_tool_summary = None
self._live_tool_stats = {}
# 检查渠道是否支持消息编辑,不支持则仅收集 token 到 buffer不实时推送
if not self._can_stream():
@@ -338,6 +364,8 @@ class StreamingHandler:
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._tool_summaries = set()
self._live_tool_summary = None
self._live_tool_stats = {}
if all_sent:
# 所有内容已通过流式发送,清空缓冲区
self._buffer = ""
@@ -382,6 +410,15 @@ class StreamingHandler:
for target_value in target_values:
bucket["targets"].add(str(target_value))
self._on_tool_stats_recorded()
def _on_tool_stats_recorded(self) -> None:
"""工具统计新增后,在支持编辑的流式渠道中立即刷新当前摘要块。"""
if not self._streaming_enabled or not self._can_stream():
return
with self._lock:
self._flush_live_tool_summary_locked()
@staticmethod
def _extract_subagent_targets(tool_kwargs: dict[str, Any]) -> list[str]:
"""按实际委派条目数量生成通用子代理统计目标。"""
@@ -404,11 +441,58 @@ class StreamingHandler:
将待输出的工具统计摘要补入缓冲区,并返回本次新增的摘要文本。
"""
with self._lock:
if self._live_tool_summary and self._pending_tool_stats:
return self._flush_live_tool_summary_locked()
summary = self._consume_pending_tool_summary_locked()
if summary:
self._buffer += summary
self._live_tool_summary = None
self._live_tool_stats = {}
return summary
def _flush_live_tool_summary_locked(self) -> str:
"""在缓冲区末尾追加或原地替换当前工具批次摘要。"""
live_summary = self._live_tool_summary
visible_buffer = self._buffer
live_start = len(self._buffer)
old_summary = None
if live_summary and self._buffer[live_summary[0] :] == live_summary[1]:
live_start, _old_block, old_summary = live_summary
visible_buffer = self._buffer[:live_start]
if self._pending_tool_stats:
self._merge_tool_stats_locked(self._live_tool_stats, self._pending_tool_stats)
self._pending_tool_stats = {}
pending_summary = self._build_tool_summary_locked(self._live_tool_stats, visible_buffer)
if not pending_summary:
return ""
summary, summary_block = pending_summary
if old_summary:
self._tool_summaries.discard(old_summary)
self._tool_summaries.add(summary)
self._buffer = visible_buffer + summary_block
self._live_tool_summary = (len(visible_buffer), summary_block, summary)
return summary_block
@staticmethod
def _merge_tool_stats_locked(
accumulated_stats: dict[str, dict[str, Any]],
new_stats: dict[str, dict[str, Any]],
) -> None:
"""将新一轮工具统计并入当前实时摘要批次。"""
for category, bucket in new_stats.items():
accumulated_bucket = accumulated_stats.setdefault(
category,
{
"count": 0,
"targets": set(),
},
)
accumulated_bucket["count"] += bucket["count"]
accumulated_bucket["targets"].update(bucket["targets"])
@staticmethod
def _classify_tool_call(
tool_name: str,
@@ -516,11 +600,36 @@ class StreamingHandler:
return "tool", None
def _consume_pending_tool_summary_locked(self) -> str:
if not self._pending_tool_stats:
pending_summary = self._build_pending_tool_summary_locked(self._buffer)
if not pending_summary:
return ""
summary, summary_block = pending_summary
self._tool_summaries.add(summary)
return summary_block
def _build_pending_tool_summary_locked(
self,
visible_buffer: str,
) -> Optional[tuple[str, str]]:
"""消费待统计数据,并按给定可见文本计算摘要块的段落边界。"""
if not self._pending_tool_stats:
return None
pending_stats = self._pending_tool_stats
self._pending_tool_stats = {}
return self._build_tool_summary_locked(pending_stats, visible_buffer)
def _build_tool_summary_locked(
self,
tool_stats: dict[str, dict[str, Any]],
visible_buffer: str,
) -> Optional[tuple[str, str]]:
"""按给定工具统计和可见文本生成摘要行及其段落块。"""
if not tool_stats:
return None
parts = []
for category, bucket in self._pending_tool_stats.items():
for category, bucket in tool_stats.items():
value = bucket["count"]
if category in {"file_read", "file_write", "directory", "web_browse", "skill"} and bucket["targets"]:
value = len(bucket["targets"])
@@ -528,20 +637,18 @@ class StreamingHandler:
if part:
parts.append(part)
self._pending_tool_stats = {}
if not parts:
return ""
return None
summary = f"{''.join(parts)}"
self._tool_summaries.add(summary)
# 摘要前始终保证一个空行,让工具执行信息与正文分属不同段落,
# 避免 Markdown 富文本把单个换行折叠成同一段落内的软换行
visible_buffer = self._buffer.rstrip(" \t")
visible_buffer = visible_buffer.rstrip(" \t")
trailing_newlines = len(visible_buffer) - len(visible_buffer.rstrip("\n"))
prefix = ""
if visible_buffer.strip():
prefix = "\n" * max(2 - trailing_newlines, 0)
return f"{prefix}{summary}\n\n"
return summary, f"{prefix}{summary}\n\n"
@staticmethod
def _format_tool_stat(category: str, count: int) -> str:

View File

@@ -121,8 +121,16 @@ class _WebAgentStreamingHandlerMixin:
tool_message=tool_message,
tool_kwargs=tool_kwargs,
)
# 结构化回调存在但未启用详细模式时,等下一段正文或流结束后一次性输出
# 多种工具的计数,避免每个调用各自生成一条摘要。
def _on_tool_stats_recorded(self) -> None:
"""WebAgent 非啰嗦模式在每次工具开始时立即发布最新摘要。"""
if self._uses_structured_tool_events():
return
if self._streaming_enabled:
self.flush_pending_tool_summary()
return
# 兼容尚未进入流式生命周期的直接调用;正式 Web 请求会在上面的
# 分支中逐次回调 SSE而不会等正文或流结束后才显示统计。
if not self._on_tool_event:
self.flush_pending_tool_summary()
@@ -177,6 +185,7 @@ class _WebAgentStreamingHandlerMixin:
self._message_response = None
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._live_tool_summary = None
async def stop_streaming(self) -> tuple[bool, str]:
"""停止 Web SSE 流式状态,保留缓冲区给 Agent 收口逻辑去重。"""
@@ -189,6 +198,7 @@ class _WebAgentStreamingHandlerMixin:
self._message_response = None
self._msg_start_offset = 0
self._pending_tool_stats = {}
self._live_tool_summary = None
return False, ""
@property

View File

@@ -3,6 +3,7 @@ from pathlib import Path
from unittest.mock import AsyncMock, patch
import langchain.agents as langchain_agents
import pytest
if not hasattr(langchain_agents, "create_agent"):
langchain_agents.create_agent = lambda *args, **kwargs: None
@@ -580,8 +581,8 @@ class TestAgentToolStreaming:
{"type": "tool", "status": "done", "tool_id": tool_id},
]
def test_web_streaming_handler_aggregates_tools_when_not_verbose(self):
"""WebAgent 关闭啰嗦模式时应隐藏逐条事件并汇总各类工具次数"""
def test_web_streaming_handler_updates_tool_summary_when_not_verbose(self):
"""WebAgent 关闭啰嗦模式时应逐次推送最新的工具计数摘要"""
async def _run():
emitted = []
@@ -614,10 +615,14 @@ class TestAgentToolStreaming:
emitted, tool_events = asyncio.run(_run())
assert tool_events == []
assert emitted == ["(执行了 1 次搜索,读取了 1 个文件)\n\n查询完成"]
assert emitted == [
"(执行了 1 次搜索)\n\n",
"(读取了 1 个文件)\n\n",
"查询完成",
]
def test_web_streaming_handler_aggregates_normal_tool_when_not_verbose(self):
"""普通工具在 WebAgent 非啰嗦模式下也应走计数汇总路径"""
def test_web_streaming_handler_updates_normal_tool_when_not_verbose(self):
"""普通工具在 WebAgent 非啰嗦模式下也应立即推送计数摘要"""
async def _run():
emitted = []
@@ -642,7 +647,66 @@ class TestAgentToolStreaming:
emitted, tool_events = asyncio.run(_run())
assert tool_events == []
assert emitted == ["(调用了 1 次工具)\n\n查询完成"]
assert emitted == ["(调用了 1 次工具)\n\n", "查询完成"]
@pytest.mark.parametrize(
"channel",
[
NotificationChannel.Telegram,
NotificationChannel.Feishu,
NotificationChannel.Slack,
NotificationChannel.Discord,
],
)
def test_editable_channels_update_tool_summary_in_place(self, channel):
"""所有支持消息编辑的通知渠道都应在原消息上更新工具统计。"""
handler = StreamingHandler()
handler._channel = channel.value
handler._source = f"{channel.value}-source"
handler._streaming_enabled = True
handler.emit("正在处理")
handler.record_tool_call(
tool_name="search_web",
tool_message="搜索网络内容",
tool_kwargs={"query": "MoviePilot"},
)
first_text = "正在处理\n\n(执行了 1 次搜索)\n\n"
assert handler._buffer == first_text
with patch("app.agent.callback.run_in_threadpool", new_callable=AsyncMock) as run_in_threadpool_mock:
run_in_threadpool_mock.return_value = MessageResponse(
message_id=f"{channel.value}-message",
chat_id=f"{channel.value}-chat",
channel=channel,
source=handler._source,
success=True,
)
asyncio.run(handler._flush())
handler.record_tool_call(
tool_name="read_file",
tool_message="读取文件",
tool_kwargs={"file_path": "/tmp/app.py"},
)
updated_text = "正在处理\n\n(执行了 1 次搜索,读取了 1 个文件)\n\n"
assert handler._buffer == updated_text
with patch("app.agent.callback.run_in_threadpool", new_callable=AsyncMock) as run_in_threadpool_mock:
run_in_threadpool_mock.return_value = True
asyncio.run(handler._flush())
assert run_in_threadpool_mock.await_count == 1
assert run_in_threadpool_mock.await_args.args[0].__name__ == "edit_message"
assert run_in_threadpool_mock.await_args.kwargs["text"] == updated_text
assert handler._sent_text == updated_text
handler.emit("第一批完成")
handler.record_tool_call(
tool_name="search_web",
tool_message="搜索下一批内容",
tool_kwargs={"query": "next"},
)
assert handler._buffer == f"{updated_text}第一批完成\n\n(执行了 1 次搜索)\n\n"
def test_rich_message_keeps_body_text_unquoted_for_telegram(self):
"""校验 Telegram 富文本只转换工具摘要行,正文保持原样。"""