diff --git a/app/api/endpoints/agent.py b/app/api/endpoints/agent.py index 5b889f5db..bc45f93a9 100644 --- a/app/api/endpoints/agent.py +++ b/app/api/endpoints/agent.py @@ -7,6 +7,7 @@ import shutil import subprocess import time import uuid +from collections import deque from queue import Empty, Queue from pathlib import Path 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_TRADITIONAL_IDLE_TIMEOUT_SECONDS = 2.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_NOTICE_QUEUES: dict[str, list[Queue[schemas.Notification]]] = {} _WEB_AGENT_NOTICE_LOCK = Lock() @@ -57,6 +62,107 @@ _WEB_AGENT_NOTICE_LISTENER_REGISTERED = False _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: """校验当前用户是否为超级管理员。""" if not getattr(user, "is_superuser", False): @@ -268,6 +374,18 @@ class _WebAgentMoviePilotAgent(MoviePilotAgent): """文本输出交由 Web 流式处理器统一回调,避免重复增量。""" 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: """ @@ -1791,14 +1909,52 @@ async def web_agent_stream( {"session_id": session_id}, locale=locale, ) - events = await _collect_web_agent_traditional_events( - text=prompt, - current_user=current_user, - original_message_id=payload.original_message_id, - original_chat_id=payload.original_chat_id, + collection_task = asyncio.create_task( + _collect_web_agent_traditional_events( + text=prompt, + current_user=current_user, + 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) 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: event_payload = copy.deepcopy(event) yield _build_web_agent_sse( @@ -1807,21 +1963,14 @@ async def web_agent_stream( locale=locale, ) if await request.is_disconnected(): - break - 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, - ) + return yield _build_web_agent_sse("done", {}, locale=locale) return StreamingResponse( traditional_event_generator(), media_type="text/event-stream", headers={ - "Cache-Control": "no-cache", + "Cache-Control": "no-cache, no-transform", "Connection": "keep-alive", "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) MessageChain().bind_user_session(str(current_user.id), session_id) - event_queue: asyncio.Queue = asyncio.Queue() - last_output = "" + event_publisher = _WebAgentEventPublisher() user_attachments = _build_web_agent_input_attachments( images=payload.images or [], files=[ @@ -1894,16 +2042,13 @@ async def web_agent_stream( ) 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): _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: """ @@ -1911,7 +2056,7 @@ async def web_agent_stream( """ for item in _build_web_agent_notification_events(notification): _apply_web_agent_display_event(item, assistant_display_message) - event_queue.put_nowait(item) + event_publisher.publish(item) async def event_generator() -> AsyncIterator[str]: """ @@ -1953,10 +2098,12 @@ async def web_agent_stream( "message": f"智能助手执行失败: {str(err)}", } _apply_web_agent_display_event(error_event, assistant_display_message) - await event_queue.put(error_event) + event_publisher.publish(error_event) finally: done_event = {"type": "done"} _apply_web_agent_display_event(done_event, assistant_display_message) + # 终态先进入 SSE 队列,避免展示快照落库延迟前端结束动画。 + event_publisher.publish(done_event) await run_in_threadpool( _save_web_agent_display_snapshot, session_id=session_id, @@ -1964,40 +2111,51 @@ async def web_agent_stream( messages=display_messages, client_session_id=payload.session_id or session_id, ) - await event_queue.put(done_event) task = asyncio.create_task(run_agent()) _WEB_AGENT_BACKGROUND_TASKS.add(task) task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard) + disconnected = False + terminal_sent = False try: yield _build_web_agent_sse( "start", {"session_id": session_id}, locale=locale, ) - disconnected = False while not global_vars.is_system_stopped: if await request.is_disconnected(): disconnected = True 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( - event.pop("type"), - event, + event_type, + {key: value for key, value in event.items() if key != "type"}, locale=locale, ) - if task.done() and event_queue.empty(): + if event_type == "done": break except asyncio.CancelledError: disconnected = True return finally: - if not task.done() and not disconnected: + if not task.done() and not disconnected and not terminal_sent: task.cancel() try: await task except asyncio.CancelledError: pass + await event_publisher.aclose() # 客户端退到后台导致 SSE 断开时,保留后台 Agent 继续执行;完成后会保存展示快照, # 前端恢复可见时可通过会话详情接口拉取最终状态。 @@ -2005,7 +2163,7 @@ async def web_agent_stream( event_generator(), media_type="text/event-stream", headers={ - "Cache-Control": "no-cache", + "Cache-Control": "no-cache, no-transform", "Connection": "keep-alive", "X-Accel-Buffering": "no", }, diff --git a/docker/nginx.common.conf b/docker/nginx.common.conf index faf6f793f..a9ca94c26 100644 --- a/docker/nginx.common.conf +++ b/docker/nginx.common.conf @@ -31,15 +31,16 @@ location /cookiecloud { } # 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; # SSE MIME类型设置 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; + gzip off; proxy_buffering off; proxy_cache off; diff --git a/tests/test_web_agent_stream.py b/tests/test_web_agent_stream.py index 9e819a4c4..dda1b8e86 100644 --- a/tests/test_web_agent_stream.py +++ b/tests/test_web_agent_stream.py @@ -1,6 +1,7 @@ import asyncio import time from queue import Queue +from threading import Event as ThreadEvent from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -8,6 +9,7 @@ from app import schemas from app.agent import ReplyMode, agent_manager from app.api.endpoints.agent import ( _WebAgentMoviePilotAgent, + _WebAgentEventPublisher, _WEB_AGENT_FILE_REGISTRY, _WEB_AGENT_NOTICE_QUEUES, _apply_web_agent_display_event, @@ -76,6 +78,32 @@ def test_split_web_agent_output_preserves_standalone_newline_delta(): 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(): """同一用户和前端会话标识应生成稳定的服务端会话 ID。""" 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 +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(): """WebAgent 渠道应声明流式、多媒体和文件发送能力。""" assert ChannelCapabilityManager.supports_capability( @@ -623,13 +671,16 @@ def test_web_agent_stream_binds_session_to_agent_manager(): if worker: worker.cancel() + async def scenario(): + response = await web_agent_stream(payload, request, user) + return "".join(await _collect_streaming_response(response)) + try: with patch("app.api.endpoints.agent.settings.AI_AGENT_ENABLE", True), patch( "app.api.endpoints.agent._WebAgentMoviePilotAgent", FakeWebAgent, ): - response = asyncio.run(web_agent_stream(payload, request, user)) - body = "".join(asyncio.run(_collect_streaming_response(response))) + body = asyncio.run(scenario()) assert "状态正常" in body assert MessageChain._user_sessions["1"][0] == session_id @@ -645,6 +696,197 @@ def test_web_agent_stream_binds_session_to_agent_manager(): 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): """读取 StreamingResponse,便于断言 SSE 内容。""" chunks = []