perf(agent): optimize web SSE streaming

This commit is contained in:
jxxghp
2026-08-07 12:45:36 +08:00
parent 759b9e47eb
commit 865635c59d
3 changed files with 437 additions and 36 deletions
+190 -32
View File
@@ -7,6 +7,7 @@ import shutil
import subprocess import subprocess
import time import time
import uuid import uuid
from collections import deque
from queue import Empty, Queue from queue import Empty, Queue
from pathlib import Path from pathlib import Path
from threading import Lock from threading import Lock
@@ -50,6 +51,10 @@ WEB_AGENT_UPLOAD_CHUNK_SIZE = 1024 * 1024
WEB_AGENT_BROWSER_AUDIO_SUFFIXES = {".aac", ".m4a", ".mp3", ".mp4", ".wav", ".wave"} WEB_AGENT_BROWSER_AUDIO_SUFFIXES = {".aac", ".m4a", ".mp3", ".mp4", ".wav", ".wave"}
WEB_AGENT_TRADITIONAL_IDLE_TIMEOUT_SECONDS = 2.0 WEB_AGENT_TRADITIONAL_IDLE_TIMEOUT_SECONDS = 2.0
WEB_AGENT_TRADITIONAL_MAX_WAIT_SECONDS = 60.0 WEB_AGENT_TRADITIONAL_MAX_WAIT_SECONDS = 60.0
WEB_AGENT_STREAM_COALESCE_SECONDS = 0.03
WEB_AGENT_STREAM_COALESCE_MAX_CHARS = 256
WEB_AGENT_STREAM_HEARTBEAT_SECONDS = 15.0
WEB_AGENT_STREAM_QUEUE_MAX_SIZE = 64
_WEB_AGENT_FILE_REGISTRY: dict[str, dict[str, Any]] = {} _WEB_AGENT_FILE_REGISTRY: dict[str, dict[str, Any]] = {}
_WEB_AGENT_NOTICE_QUEUES: dict[str, list[Queue[schemas.Notification]]] = {} _WEB_AGENT_NOTICE_QUEUES: dict[str, list[Queue[schemas.Notification]]] = {}
_WEB_AGENT_NOTICE_LOCK = Lock() _WEB_AGENT_NOTICE_LOCK = Lock()
@@ -57,6 +62,107 @@ _WEB_AGENT_NOTICE_LISTENER_REGISTERED = False
_WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task] = set() _WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task] = set()
class _WebAgentEventPublisher:
"""合并 WebAgent 文本增量,并通过有界队列向 SSE 消费者提供事件。"""
def __init__(self) -> None:
self._queue: asyncio.Queue[dict] = asyncio.Queue(
maxsize=WEB_AGENT_STREAM_QUEUE_MAX_SIZE
)
self._pending_events: deque[dict] = deque()
self._pending_signal = asyncio.Event()
self._pending_delta = ""
self._delta_timer: Optional[asyncio.TimerHandle] = None
self._disposed = False
self._max_depth = 0
self._last_logged_depth = 0
self._pump_task = asyncio.create_task(self._pump())
@property
def max_depth(self) -> int:
"""返回本轮发布器观测到的最大积压深度。"""
return self._max_depth
def publish(self, event: dict) -> None:
"""发布事件;相邻文本会按时间或长度边界合并。"""
if self._disposed:
return
if event.get("type") == "delta":
self._pending_delta += str(event.get("content") or "")
if len(self._pending_delta) >= WEB_AGENT_STREAM_COALESCE_MAX_CHARS:
self._flush_delta()
elif self._delta_timer is None:
loop = asyncio.get_running_loop()
self._delta_timer = loop.call_later(
WEB_AGENT_STREAM_COALESCE_SECONDS,
self._flush_delta,
)
return
self._flush_delta()
self._append_event(event)
async def get(self) -> dict:
"""等待并返回下一条已排序事件。"""
return await self._queue.get()
async def aclose(self) -> None:
"""停止发布器并释放等待中的泵任务。"""
if self._disposed:
return
self._disposed = True
self._cancel_delta_timer()
self._pending_delta = ""
self._pending_events.clear()
self._pump_task.cancel()
try:
await self._pump_task
except asyncio.CancelledError:
pass
def _cancel_delta_timer(self) -> None:
"""取消尚未触发的文本合并计时器。"""
if self._delta_timer is None:
return
self._delta_timer.cancel()
self._delta_timer = None
def _flush_delta(self) -> None:
"""把当前文本缓冲转换成一条增量事件。"""
self._cancel_delta_timer()
if not self._pending_delta or self._disposed:
return
content = self._pending_delta
self._pending_delta = ""
self._append_event({"type": "delta", "content": content})
def _append_event(self, event: dict) -> None:
"""追加待发布事件,相邻文本在出口阻塞时继续合并。"""
if (
event.get("type") == "delta"
and self._pending_events
and self._pending_events[-1].get("type") == "delta"
):
self._pending_events[-1]["content"] += str(event.get("content") or "")
else:
self._pending_events.append(event)
self._pending_signal.set()
depth = self._queue.qsize() + len(self._pending_events)
self._max_depth = max(self._max_depth, depth)
if depth >= WEB_AGENT_STREAM_QUEUE_MAX_SIZE // 2 and depth > self._last_logged_depth:
self._last_logged_depth = depth
logger.debug(f"WebAgent SSE事件积压深度: {depth}")
async def _pump(self) -> None:
"""按发布顺序把本地合并结果写入有界出口队列。"""
while True:
await self._pending_signal.wait()
while self._pending_events:
event = self._pending_events.popleft()
await self._queue.put(event)
self._pending_signal.clear()
def _ensure_superuser(user: User) -> None: def _ensure_superuser(user: User) -> None:
"""校验当前用户是否为超级管理员。""" """校验当前用户是否为超级管理员。"""
if not getattr(user, "is_superuser", False): if not getattr(user, "is_superuser", False):
@@ -268,6 +374,18 @@ class _WebAgentMoviePilotAgent(MoviePilotAgent):
"""文本输出交由 Web 流式处理器统一回调,避免重复增量。""" """文本输出交由 Web 流式处理器统一回调,避免重复增量。"""
self.stream_handler.emit(text) self.stream_handler.emit(text)
def _emit_output(self, text: str) -> None:
"""保留完整输出状态,同时只把本次增量交给 Web SSE 回调。"""
if not text:
return
self._streamed_output += text
if not callable(self.output_callback):
return
try:
self.output_callback(text)
except Exception as e:
logger.debug(f"Web智能体输出回调失败: {e}")
def _build_web_agent_session_id(user: User, session_id: Optional[str]) -> str: def _build_web_agent_session_id(user: User, session_id: Optional[str]) -> str:
""" """
@@ -1791,14 +1909,52 @@ async def web_agent_stream(
{"session_id": session_id}, {"session_id": session_id},
locale=locale, locale=locale,
) )
events = await _collect_web_agent_traditional_events( collection_task = asyncio.create_task(
text=prompt, _collect_web_agent_traditional_events(
current_user=current_user, text=prompt,
original_message_id=payload.original_message_id, current_user=current_user,
original_chat_id=payload.original_chat_id, original_message_id=payload.original_message_id,
original_chat_id=payload.original_chat_id,
)
) )
try:
while True:
try:
events = await asyncio.wait_for(
asyncio.shield(collection_task),
timeout=WEB_AGENT_STREAM_HEARTBEAT_SECONDS,
)
break
except asyncio.TimeoutError:
if await request.is_disconnected():
collection_task.cancel()
return
yield ": heartbeat\n\n"
except asyncio.CancelledError:
if not collection_task.done():
collection_task.cancel()
return
assistant_message = _build_web_agent_display_message_from_events(events) assistant_message = _build_web_agent_display_message_from_events(events)
display_messages.append(assistant_message) display_messages.append(assistant_message)
async def save_display_snapshot() -> None:
"""后台保存传统消息展示快照,不阻塞 SSE 终态。"""
try:
await run_in_threadpool(
_save_web_agent_display_snapshot,
session_id=session_id,
current_user=current_user,
messages=display_messages,
client_session_id=payload.session_id or session_id,
)
except Exception as err:
logger.error(f"保存WebAgent传统消息快照失败: {str(err)}")
snapshot_task = asyncio.create_task(save_display_snapshot())
_WEB_AGENT_BACKGROUND_TASKS.add(snapshot_task)
snapshot_task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
await asyncio.sleep(0)
for event in events: for event in events:
event_payload = copy.deepcopy(event) event_payload = copy.deepcopy(event)
yield _build_web_agent_sse( yield _build_web_agent_sse(
@@ -1807,21 +1963,14 @@ async def web_agent_stream(
locale=locale, locale=locale,
) )
if await request.is_disconnected(): if await request.is_disconnected():
break return
await run_in_threadpool(
_save_web_agent_display_snapshot,
session_id=session_id,
current_user=current_user,
messages=display_messages,
client_session_id=payload.session_id or session_id,
)
yield _build_web_agent_sse("done", {}, locale=locale) yield _build_web_agent_sse("done", {}, locale=locale)
return StreamingResponse( return StreamingResponse(
traditional_event_generator(), traditional_event_generator(),
media_type="text/event-stream", media_type="text/event-stream",
headers={ headers={
"Cache-Control": "no-cache", "Cache-Control": "no-cache, no-transform",
"Connection": "keep-alive", "Connection": "keep-alive",
"X-Accel-Buffering": "no", "X-Accel-Buffering": "no",
}, },
@@ -1868,8 +2017,7 @@ async def web_agent_stream(
session_id = _build_web_agent_session_id(current_user, payload.session_id) session_id = _build_web_agent_session_id(current_user, payload.session_id)
MessageChain().bind_user_session(str(current_user.id), session_id) MessageChain().bind_user_session(str(current_user.id), session_id)
event_queue: asyncio.Queue = asyncio.Queue() event_publisher = _WebAgentEventPublisher()
last_output = ""
user_attachments = _build_web_agent_input_attachments( user_attachments = _build_web_agent_input_attachments(
images=payload.images or [], images=payload.images or [],
files=[ files=[
@@ -1894,16 +2042,13 @@ async def web_agent_stream(
) )
display_messages.append(assistant_display_message) display_messages.append(assistant_display_message)
def output_callback(output: str) -> None: def output_callback(delta: str) -> None:
""" """
接收 Agent 累积输出并转成增量事件。 接收 Agent 文本增量并转换成 SSE 事件。
""" """
nonlocal last_output
delta = output[len(last_output):] if output.startswith(last_output) else output
last_output = output
for item in _split_web_agent_output(delta): for item in _split_web_agent_output(delta):
_apply_web_agent_display_event(item, assistant_display_message) _apply_web_agent_display_event(item, assistant_display_message)
event_queue.put_nowait(item) event_publisher.publish(item)
def notification_callback(notification: schemas.Notification) -> None: def notification_callback(notification: schemas.Notification) -> None:
""" """
@@ -1911,7 +2056,7 @@ async def web_agent_stream(
""" """
for item in _build_web_agent_notification_events(notification): for item in _build_web_agent_notification_events(notification):
_apply_web_agent_display_event(item, assistant_display_message) _apply_web_agent_display_event(item, assistant_display_message)
event_queue.put_nowait(item) event_publisher.publish(item)
async def event_generator() -> AsyncIterator[str]: async def event_generator() -> AsyncIterator[str]:
""" """
@@ -1953,10 +2098,12 @@ async def web_agent_stream(
"message": f"智能助手执行失败: {str(err)}", "message": f"智能助手执行失败: {str(err)}",
} }
_apply_web_agent_display_event(error_event, assistant_display_message) _apply_web_agent_display_event(error_event, assistant_display_message)
await event_queue.put(error_event) event_publisher.publish(error_event)
finally: finally:
done_event = {"type": "done"} done_event = {"type": "done"}
_apply_web_agent_display_event(done_event, assistant_display_message) _apply_web_agent_display_event(done_event, assistant_display_message)
# 终态先进入 SSE 队列,避免展示快照落库延迟前端结束动画。
event_publisher.publish(done_event)
await run_in_threadpool( await run_in_threadpool(
_save_web_agent_display_snapshot, _save_web_agent_display_snapshot,
session_id=session_id, session_id=session_id,
@@ -1964,40 +2111,51 @@ async def web_agent_stream(
messages=display_messages, messages=display_messages,
client_session_id=payload.session_id or session_id, client_session_id=payload.session_id or session_id,
) )
await event_queue.put(done_event)
task = asyncio.create_task(run_agent()) task = asyncio.create_task(run_agent())
_WEB_AGENT_BACKGROUND_TASKS.add(task) _WEB_AGENT_BACKGROUND_TASKS.add(task)
task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard) task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
disconnected = False
terminal_sent = False
try: try:
yield _build_web_agent_sse( yield _build_web_agent_sse(
"start", "start",
{"session_id": session_id}, {"session_id": session_id},
locale=locale, locale=locale,
) )
disconnected = False
while not global_vars.is_system_stopped: while not global_vars.is_system_stopped:
if await request.is_disconnected(): if await request.is_disconnected():
disconnected = True disconnected = True
break break
event = await event_queue.get() try:
event = await asyncio.wait_for(
event_publisher.get(),
timeout=WEB_AGENT_STREAM_HEARTBEAT_SECONDS,
)
except asyncio.TimeoutError:
yield ": heartbeat\n\n"
continue
event_type = str(event.get("type") or "")
if event_type == "done":
terminal_sent = True
yield _build_web_agent_sse( yield _build_web_agent_sse(
event.pop("type"), event_type,
event, {key: value for key, value in event.items() if key != "type"},
locale=locale, locale=locale,
) )
if task.done() and event_queue.empty(): if event_type == "done":
break break
except asyncio.CancelledError: except asyncio.CancelledError:
disconnected = True disconnected = True
return return
finally: finally:
if not task.done() and not disconnected: if not task.done() and not disconnected and not terminal_sent:
task.cancel() task.cancel()
try: try:
await task await task
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
await event_publisher.aclose()
# 客户端退到后台导致 SSE 断开时,保留后台 Agent 继续执行;完成后会保存展示快照, # 客户端退到后台导致 SSE 断开时,保留后台 Agent 继续执行;完成后会保存展示快照,
# 前端恢复可见时可通过会话详情接口拉取最终状态。 # 前端恢复可见时可通过会话详情接口拉取最终状态。
@@ -2005,7 +2163,7 @@ async def web_agent_stream(
event_generator(), event_generator(),
media_type="text/event-stream", media_type="text/event-stream",
headers={ headers={
"Cache-Control": "no-cache", "Cache-Control": "no-cache, no-transform",
"Connection": "keep-alive", "Connection": "keep-alive",
"X-Accel-Buffering": "no", "X-Accel-Buffering": "no",
}, },
+3 -2
View File
@@ -31,15 +31,16 @@ location /cookiecloud {
} }
# SSE特殊配置 # SSE特殊配置
location ~ ^/api/v1/(system/(message|progress/|logging)|search/.*/stream$) { location ~ ^/api/v1/(system/(message|progress/|logging)|search/.*/stream$|message/agent/stream$) {
error_log /dev/null crit; error_log /dev/null crit;
# SSE MIME类型设置 # SSE MIME类型设置
default_type text/event-stream; default_type text/event-stream;
# 禁用缓存 # 禁用缓存
add_header Cache-Control no-cache; add_header Cache-Control "no-cache, no-transform";
add_header X-Accel-Buffering no; add_header X-Accel-Buffering no;
gzip off;
proxy_buffering off; proxy_buffering off;
proxy_cache off; proxy_cache off;
+244 -2
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import time import time
from queue import Queue from queue import Queue
from threading import Event as ThreadEvent
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
@@ -8,6 +9,7 @@ from app import schemas
from app.agent import ReplyMode, agent_manager from app.agent import ReplyMode, agent_manager
from app.api.endpoints.agent import ( from app.api.endpoints.agent import (
_WebAgentMoviePilotAgent, _WebAgentMoviePilotAgent,
_WebAgentEventPublisher,
_WEB_AGENT_FILE_REGISTRY, _WEB_AGENT_FILE_REGISTRY,
_WEB_AGENT_NOTICE_QUEUES, _WEB_AGENT_NOTICE_QUEUES,
_apply_web_agent_display_event, _apply_web_agent_display_event,
@@ -76,6 +78,32 @@ def test_split_web_agent_output_preserves_standalone_newline_delta():
assert content == "可以这样操作:\n- **搜索资源**:搜索电影\n- **下载管理**:添加任务" assert content == "可以这样操作:\n- **搜索资源**:搜索电影\n- **下载管理**:添加任务"
def test_web_agent_event_publisher_coalesces_text_before_semantic_events():
"""连续文本应合并,且工具事件前的文本顺序不能改变。"""
async def scenario():
publisher = _WebAgentEventPublisher()
try:
for index in range(100):
publisher.publish({"type": "delta", "content": str(index % 10)})
publisher.publish({"type": "tool", "message": "查询完成"})
first = await asyncio.wait_for(publisher.get(), timeout=1)
second = await asyncio.wait_for(publisher.get(), timeout=1)
return first, second, publisher.max_depth
finally:
await publisher.aclose()
first, second, max_depth = asyncio.run(scenario())
assert first == {
"type": "delta",
"content": "".join(str(index % 10) for index in range(100)),
}
assert second == {"type": "tool", "message": "查询完成"}
assert max_depth == 2
def test_build_web_agent_session_id_is_stable_per_user_and_seed(): def test_build_web_agent_session_id_is_stable_per_user_and_seed():
"""同一用户和前端会话标识应生成稳定的服务端会话 ID。""" """同一用户和前端会话标识应生成稳定的服务端会话 ID。"""
user = SimpleNamespace(id=1, name="admin") user = SimpleNamespace(id=1, name="admin")
@@ -347,6 +375,26 @@ def test_web_agent_reused_for_background_task_disables_streaming():
assert agent._should_stream() is False assert agent._should_stream() is False
def test_web_agent_output_callback_receives_only_new_text():
"""WebAgent 外部回调应接收增量,同时内部仍保留完整输出。"""
outputs = []
agent = _WebAgentMoviePilotAgent(
session_id="web-agent:incremental-output",
user_id="7",
channel=MessageChannel.WebAgent.value,
source="web-agent",
username="admin",
replay_mode=ReplyMode.CAPTURE_ONLY,
output_callback=outputs.append,
)
agent._handle_stream_text("")
agent._handle_stream_text("")
assert outputs == ["", ""]
assert agent._streamed_output == "你好"
def test_web_agent_channel_supports_streaming_and_attachments(): def test_web_agent_channel_supports_streaming_and_attachments():
"""WebAgent 渠道应声明流式、多媒体和文件发送能力。""" """WebAgent 渠道应声明流式、多媒体和文件发送能力。"""
assert ChannelCapabilityManager.supports_capability( assert ChannelCapabilityManager.supports_capability(
@@ -623,13 +671,16 @@ def test_web_agent_stream_binds_session_to_agent_manager():
if worker: if worker:
worker.cancel() worker.cancel()
async def scenario():
response = await web_agent_stream(payload, request, user)
return "".join(await _collect_streaming_response(response))
try: try:
with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch( with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch(
"app.api.endpoints.agent._WebAgentMoviePilotAgent", "app.api.endpoints.agent._WebAgentMoviePilotAgent",
FakeWebAgent, FakeWebAgent,
): ):
response = asyncio.run(web_agent_stream(payload, request, user)) body = asyncio.run(scenario())
body = "".join(asyncio.run(_collect_streaming_response(response)))
assert "状态正常" in body assert "状态正常" in body
assert MessageChain._user_sessions["1"][0] == session_id assert MessageChain._user_sessions["1"][0] == session_id
@@ -645,6 +696,197 @@ def test_web_agent_stream_binds_session_to_agent_manager():
worker.cancel() worker.cancel()
def test_web_agent_stream_emits_heartbeat_during_idle_tool_wait():
"""长时间没有 Agent 事件时应发送 SSE heartbeat 保持连接。"""
payload = schemas.AgentWebChatRequest(text="分析系统状态", session_id="browser-heartbeat")
request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False))
user = SimpleNamespace(id=1, name="admin", is_superuser=True)
async def slow_process_message(**kwargs):
"""模拟工具执行期间暂时没有可见输出。"""
await asyncio.sleep(0.035)
kwargs["output_callback"]("状态正常")
async def scenario():
response = await web_agent_stream(payload, request, user)
return "".join(await _collect_streaming_response(response))
with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch(
"app.api.endpoints.agent.WEB_AGENT_STREAM_HEARTBEAT_SECONDS",
0.01,
), patch(
"app.api.endpoints.agent._is_web_agent_traditional_message",
return_value=False,
), patch(
"app.api.endpoints.agent._has_web_agent_traditional_interaction",
return_value=False,
), patch(
"app.api.endpoints.agent._build_web_agent_session_id",
return_value="web-agent:heartbeat",
), patch.object(
MessageChain,
"bind_user_session",
), patch.object(
agent_manager,
"process_message",
side_effect=slow_process_message,
), patch(
"app.api.endpoints.agent._save_web_agent_display_snapshot",
):
body = asyncio.run(scenario())
assert ": heartbeat\n\n" in body
assert '"type": "delta"' in body
assert '"type": "done"' in body
def test_web_agent_traditional_stream_keeps_alive_and_saves_after_done():
"""传统消息等待期间应保活,且展示快照不能阻塞终态。"""
payload = schemas.AgentWebChatRequest(text="/状态", session_id="traditional-heartbeat")
request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False))
user = SimpleNamespace(id=1, name="admin", is_superuser=True)
snapshot_started = ThreadEvent()
snapshot_release = ThreadEvent()
snapshot_finished = ThreadEvent()
async def slow_collect(**_kwargs):
"""模拟传统消息链路等待外部结果。"""
await asyncio.sleep(0.035)
return [{"type": "delta", "content": "状态正常"}]
def slow_snapshot(**_kwargs):
"""阻塞快照写入,便于断言 done 不等待落库。"""
snapshot_started.set()
snapshot_release.wait(timeout=2)
snapshot_finished.set()
async def scenario():
response = await web_agent_stream(payload, request, user)
assert response.headers["cache-control"] == "no-cache, no-transform"
iterator = response.body_iterator.__aiter__()
received = []
while True:
chunk = await asyncio.wait_for(anext(iterator), timeout=1)
text = chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
received.append(text)
if '"type": "done"' in text:
break
for _ in range(100):
if snapshot_started.is_set():
break
await asyncio.sleep(0.001)
assert snapshot_started.is_set()
assert not snapshot_finished.is_set()
await iterator.aclose()
return "".join(received)
try:
with patch(
"app.api.endpoints.agent.WEB_AGENT_STREAM_HEARTBEAT_SECONDS",
0.01,
), patch(
"app.api.endpoints.agent._is_web_agent_traditional_message",
return_value=True,
), patch(
"app.api.endpoints.agent._ensure_web_agent_command_allowed",
return_value=None,
), patch(
"app.api.endpoints.agent._get_web_agent_unknown_command_message",
return_value=None,
), patch(
"app.api.endpoints.agent._build_web_agent_session_id",
return_value="web-agent:traditional-heartbeat",
), patch(
"app.api.endpoints.agent._collect_web_agent_traditional_events",
side_effect=slow_collect,
), patch(
"app.api.endpoints.agent._save_web_agent_display_snapshot",
side_effect=slow_snapshot,
):
body = asyncio.run(scenario())
assert ": heartbeat\n\n" in body
assert '"type": "delta"' in body
assert '"type": "done"' in body
assert not snapshot_finished.is_set()
finally:
snapshot_release.set()
assert snapshot_finished.wait(timeout=1)
def test_web_agent_stream_sends_done_before_snapshot_persistence_finishes():
"""展示快照落库缓慢时,前端终态不应被数据库操作阻塞。"""
payload = schemas.AgentWebChatRequest(text="检查系统", session_id="browser-snapshot")
request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False))
user = SimpleNamespace(id=1, name="admin", is_superuser=True)
snapshot_started = ThreadEvent()
snapshot_release = ThreadEvent()
snapshot_finished = ThreadEvent()
async def immediate_process_message(**kwargs):
"""立即生成一段文本,随后进入终态。"""
kwargs["output_callback"]("检查完成")
def slow_snapshot(**_kwargs):
"""阻塞快照写入,便于验证 done 的发送时机。"""
snapshot_started.set()
snapshot_release.wait(timeout=2)
snapshot_finished.set()
async def scenario():
response = await web_agent_stream(payload, request, user)
iterator = response.body_iterator.__aiter__()
received = []
while True:
chunk = await asyncio.wait_for(anext(iterator), timeout=1)
text = chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
received.append(text)
if '"type": "done"' in text:
break
for _ in range(100):
if snapshot_started.is_set():
break
await asyncio.sleep(0.001)
assert snapshot_started.is_set()
assert not snapshot_finished.is_set()
await iterator.aclose()
return "".join(received)
try:
with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch(
"app.api.endpoints.agent._is_web_agent_traditional_message",
return_value=False,
), patch(
"app.api.endpoints.agent._has_web_agent_traditional_interaction",
return_value=False,
), patch(
"app.api.endpoints.agent._build_web_agent_session_id",
return_value="web-agent:snapshot",
), patch.object(
MessageChain,
"bind_user_session",
), patch.object(
agent_manager,
"process_message",
side_effect=immediate_process_message,
), patch(
"app.api.endpoints.agent._save_web_agent_display_snapshot",
side_effect=slow_snapshot,
):
body = asyncio.run(scenario())
assert '"type": "done"' in body
assert not snapshot_finished.is_set()
finally:
snapshot_release.set()
assert snapshot_finished.wait(timeout=1)
async def _collect_streaming_response(response): async def _collect_streaming_response(response):
"""读取 StreamingResponse,便于断言 SSE 内容。""" """读取 StreamingResponse,便于断言 SSE 内容。"""
chunks = [] chunks = []