mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 07:27:15 +08:00
perf(agent): optimize web SSE streaming
This commit is contained in:
+190
-32
@@ -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",
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|
||||||
|
|||||||
@@ -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 = []
|
||||||
|
|||||||
Reference in New Issue
Block a user