mirror of
https://hubproxy.babadafafafafa.cn/https://github.com/jxxghp/MoviePilot.git
synced 2026-09-20 08:03:34 +08:00
feat(agent): stream non-verbose tool progress
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 富文本只转换工具摘要行,正文保持原样。"""
|
||||
|
||||
Reference in New Issue
Block a user