diff --git a/backend/package/yuxi/agents/base.py b/backend/package/yuxi/agents/base.py index d136b473..7c5ba075 100644 --- a/backend/package/yuxi/agents/base.py +++ b/backend/package/yuxi/agents/base.py @@ -96,6 +96,32 @@ class BaseAgent: ): yield msg, metadata + async def stream_messages_with_state(self, messages: list[str], input_context=None, **kwargs): + context = self.context_schema() + context.update_from_dict(input_context or {}) + graph = await self.get_graph(context=context) + logger.debug(f"stream_messages_with_state: {context=}") + + input_config = { + "configurable": {"thread_id": context.thread_id, "user_id": context.user_id}, + "recursion_limit": 300, + } + + if callbacks := kwargs.get("callbacks"): + input_config["callbacks"] = list(callbacks) + if metadata := kwargs.get("metadata"): + input_config["metadata"] = dict(metadata) + if tags := kwargs.get("tags"): + input_config["tags"] = list(tags) + + async for mode, payload in graph.astream( + {"messages": messages}, + stream_mode=["messages", "values"], + context=context, + config=input_config, + ): + yield mode, payload + async def invoke_messages(self, messages: list[str], input_context=None, **kwargs): context = self.context_schema() context.update_from_dict(input_context or {}) diff --git a/backend/package/yuxi/agents/buildin/deep_agent/graph.py b/backend/package/yuxi/agents/buildin/deep_agent/graph.py index 12d1ef43..3050b9ec 100644 --- a/backend/package/yuxi/agents/buildin/deep_agent/graph.py +++ b/backend/package/yuxi/agents/buildin/deep_agent/graph.py @@ -27,7 +27,7 @@ from .prompt import DEEP_PROMPT class DeepAgent(BaseAgent): name = "深度分析" description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务" - capabilities = ["file_upload", "files", "todo"] # 支持文件上传功能 + capabilities = ["file_upload", "files"] # 支持文件上传功能 metadata = {"examples": ["调研一下多模态 GraphRAG 的相关论文"]} def __init__(self, **kwargs): diff --git a/backend/package/yuxi/services/chat_service.py b/backend/package/yuxi/services/chat_service.py index fa60c0f0..30620350 100644 --- a/backend/package/yuxi/services/chat_service.py +++ b/backend/package/yuxi/services/chat_service.py @@ -118,6 +118,29 @@ def extract_agent_state(values: dict) -> AgentStatePayload: return result +def _agent_state_signature(agent_state: AgentStatePayload | dict | None) -> str: + if not agent_state: + return "" + try: + return json.dumps(agent_state, ensure_ascii=False, sort_keys=True) + except Exception: + return str(agent_state) + + +async def _stream_agent_events(agent, messages, *, input_context=None, **kwargs): + if hasattr(agent, "stream_messages_with_state"): + async for mode, payload in agent.stream_messages_with_state( + messages, + input_context=input_context, + **kwargs, + ): + yield mode, payload + return + + async for msg, metadata in agent.stream_messages(messages, input_context=input_context, **kwargs): + yield "messages", (msg, metadata) + + async def _get_existing_message_ids(conv_repo: ConversationRepository, thread_id: str) -> set[str]: existing_messages = await conv_repo.get_messages_by_thread_id(thread_id) return { @@ -536,6 +559,7 @@ async def agent_chat( message_type=message_type, ) trace_info: dict[str, Any] = {} + last_agent_state_signature = "" try: conv_repo = ConversationRepository(db) @@ -738,6 +762,7 @@ async def stream_agent_chat( full_msg = None accumulated_content: list[str] = [] trace_info: dict[str, Any] = {} + last_agent_state_signature = "" try: conv_repo = ConversationRepository(db) @@ -769,13 +794,23 @@ async def stream_agent_chat( full_msg = None accumulated_content = [] - async for msg, metadata in agent.stream_messages( + async for mode, payload in _stream_agent_events( + agent, messages, input_context=input_context, callbacks=langfuse_run.callbacks, metadata=langfuse_run.metadata, tags=langfuse_run.tags, ): + if mode == "values": + agent_state = extract_agent_state(payload if isinstance(payload, dict) else {}) + signature = _agent_state_signature(agent_state) + if signature and signature != last_agent_state_signature: + last_agent_state_signature = signature + yield make_chunk(status="agent_state", agent_state=agent_state, meta=meta) + continue + + msg, metadata = payload if isinstance(msg, AIMessageChunk): accumulated_content.append(msg.content) trace_info = get_trace_info(langfuse_run) @@ -800,16 +835,6 @@ async def stream_agent_chat( trace_info = get_trace_info(langfuse_run) yield make_chunk(msg=msg_dict, metadata=metadata, status="loading") - try: - if msg_dict.get("type") == "tool": - graph = await agent.get_graph() - state = await graph.aget_state(langgraph_config) - agent_state = extract_agent_state(getattr(state, "values", {})) if state else {} - if agent_state: - yield make_chunk(status="agent_state", agent_state=agent_state, meta=meta) - except Exception as e: - logger.error(f"Error processing tool message: {e}") - full_msg = _ensure_full_msg(full_msg, accumulated_content) trace_info = get_trace_info(langfuse_run) @@ -836,7 +861,9 @@ async def stream_agent_chat( except Exception: agent_state = {} - if agent_state: + final_signature = _agent_state_signature(agent_state) + if final_signature and final_signature != last_agent_state_signature: + last_agent_state_signature = final_signature yield make_chunk(status="agent_state", agent_state=agent_state, meta=meta) # 先存储数据库,再返回 finished,避免前端查询时数据未落库 @@ -969,11 +996,20 @@ async def stream_agent_resume( "metadata": langfuse_run.metadata, "tags": langfuse_run.tags, }, - stream_mode="messages", + stream_mode=["messages", "values"], ) try: - async for msg, metadata in stream_source: + async for mode, payload in stream_source: + if mode == "values": + agent_state = extract_agent_state(payload if isinstance(payload, dict) else {}) + signature = _agent_state_signature(agent_state) + if signature and signature != last_agent_state_signature: + last_agent_state_signature = signature + yield make_resume_chunk(status="agent_state", agent_state=agent_state, meta=meta) + continue + + msg, metadata = payload trace_info = get_trace_info(langfuse_run) msg_dict = msg.model_dump() if "id" not in msg_dict: @@ -989,6 +1025,16 @@ async def stream_agent_resume( meta["time_cost"] = asyncio.get_event_loop().time() - start_time + try: + state = await graph.aget_state(langgraph_config) + agent_state = extract_agent_state(getattr(state, "values", {})) if state else {} + except Exception: + agent_state = {} + + final_signature = _agent_state_signature(agent_state) + if final_signature and final_signature != last_agent_state_signature: + yield make_resume_chunk(status="agent_state", agent_state=agent_state, meta=meta) + # 先存储数据库,再返回 finished,避免前端查询时数据未落库 conv_repo = ConversationRepository(db) await save_messages_from_langgraph_state( diff --git a/backend/package/yuxi/services/conversation_service.py b/backend/package/yuxi/services/conversation_service.py index 46cde9f0..01e33f10 100644 --- a/backend/package/yuxi/services/conversation_service.py +++ b/backend/package/yuxi/services/conversation_service.py @@ -19,7 +19,6 @@ from yuxi.utils.paths import VIRTUAL_PATH_UPLOADS ATTACHMENT_ALLOWED_EXTENSIONS: tuple[str, ...] = () MAX_ATTACHMENT_SIZE_BYTES = 5 * 1024 * 1024 # 5 MB -MAX_ATTACHMENT_MARKDOWN_CHARS = 32_000 MAX_ATTACHMENT_MARKDOWN_CHARS = 32_000 # TODO: 转 MARKDOWN的时候,不应该裁剪 diff --git a/backend/test/unit/services/test_chat_service_langfuse_stream.py b/backend/test/unit/services/test_chat_service_langfuse_stream.py index 202d1953..d9bbf83b 100644 --- a/backend/test/unit/services/test_chat_service_langfuse_stream.py +++ b/backend/test/unit/services/test_chat_service_langfuse_stream.py @@ -150,3 +150,73 @@ async def test_stream_agent_chat_passes_langfuse_callbacks_and_persists_trace_in assert chunks[-1]["status"] == "finished" assert calls["flushed"] is True assert isinstance(calls["stream_messages"][0], HumanMessage) + + +@pytest.mark.asyncio +async def test_stream_agent_chat_emits_realtime_agent_state_from_values(monkeypatch: pytest.MonkeyPatch): + class FakeGraph: + async def aget_state(self, _config): + return SimpleNamespace(values={"todos": [{"content": "done", "status": "completed"}]}) + + class FakeAgent: + async def stream_messages_with_state(self, messages, input_context=None, **kwargs): + yield "values", {"messages": [], "todos": [{"content": "step 1", "status": "pending"}]} + yield "values", {"messages": [], "todos": [{"content": "step 1", "status": "in_progress"}]} + yield "values", {"messages": [], "todos": [{"content": "step 1", "status": "in_progress"}]} + yield "messages", (AIMessageChunk(content="hello"), {"node": "llm"}) + + async def stream_messages(self, messages, input_context=None, **kwargs): + raise AssertionError("stream_messages fallback should not be used") + + async def get_graph(self): + return FakeGraph() + + async def fake_get_agent_config_by_id(db, user, agent_config_id): + return SimpleNamespace(agent_id="test-agent", config_json={"context": {}}) + + async def fake_save_messages_from_langgraph_state(*, agent_instance, thread_id, conv_repo, config_dict, trace_info): + return None + + async def fake_guard_check(_content): + return False + + async def fake_guard_check_with_keywords(_content): + return False + + async def fake_interrupts(agent, langgraph_config, make_chunk, meta, thread_id): + if False: + yield None + return + + monkeypatch.setattr(svc.agent_manager, "get_agent", lambda agent_id: FakeAgent()) + monkeypatch.setattr(svc, "get_agent_config_by_id", fake_get_agent_config_by_id) + monkeypatch.setattr(svc, "ConversationRepository", _FakeConvRepo) + monkeypatch.setattr(svc, "save_messages_from_langgraph_state", fake_save_messages_from_langgraph_state) + monkeypatch.setattr(svc.content_guard, "check", fake_guard_check) + monkeypatch.setattr(svc.content_guard, "check_with_keywords", fake_guard_check_with_keywords) + monkeypatch.setattr(svc, "check_and_handle_interrupts", fake_interrupts) + monkeypatch.setattr( + svc, + "_build_langfuse_run_context", + lambda **kwargs: SimpleNamespace(callbacks=[], metadata={}, tags=[], trace_id=None), + ) + monkeypatch.setattr(svc, "get_trace_info", lambda _run_context: {}) + monkeypatch.setattr(svc, "flush_langfuse", lambda: None) + + chunks = [] + async for chunk in svc.stream_agent_chat( + query="hello", + agent_config_id=123, + thread_id="thread-1", + meta={"request_id": "req-1"}, + image_content=None, + current_user=SimpleNamespace(id="user-1", department_id="dept-1"), + db=object(), + ): + chunks.append(json.loads(chunk.decode("utf-8"))) + + agent_state_chunks = [chunk for chunk in chunks if chunk.get("status") == "agent_state"] + assert len(agent_state_chunks) == 3 + assert agent_state_chunks[0]["agent_state"]["todos"][0]["status"] == "pending" + assert agent_state_chunks[1]["agent_state"]["todos"][0]["status"] == "in_progress" + assert agent_state_chunks[2]["agent_state"]["todos"][0]["status"] == "completed" diff --git a/docs/develop-guides/roadmap.md b/docs/develop-guides/roadmap.md index 963875d7..ff7a13bb 100644 --- a/docs/develop-guides/roadmap.md +++ b/docs/develop-guides/roadmap.md @@ -58,6 +58,7 @@ ### 修复 +- 调整智能体 todo 展示语义:待办状态不再作为 `capabilities` 前端开关,而是直接根据运行态 `agent_state.todos` 渲染;同时将 todo 入口从 Agent Panel 移到输入框内的轻量浮层,并让右侧“状态工作台”收敛为文件系统视图,输入框按钮文案同步由“状态”调整为“文件” - 优化 Agent 输入框 mention 行为:在保留附件 mention 的同时,将共享 `workspace` 文件纳入候选范围;并将 `@` 空查询时的候选列表改为空,仅在继续输入后再执行筛选,避免工作区文件过多时直接铺满下拉面板 - 为前端工作台文件树补齐文件删除能力:`/api/viewer/filesystem/file` 新增删除接口,`AgentPanel` 文件节点新增删除按钮与确认交互,删除后会同步刷新树与预览状态 - 调整前端工作台文件预览交互:恢复默认侧边/弹窗预览,并新增显式“全屏预览”入口;全屏模式下由预览内容直接覆盖整页,仅保留右上角悬浮关闭按钮;同时修复 HTML 文件首次在弹窗中预览偶现白屏的问题,改为在内容更新后强制重建 `iframe`