feat(agent): 添加流式消息处理功能,支持实时更新智能体状态

This commit is contained in:
Wenjie Zhang 2026-04-01 03:17:43 +08:00
parent b1a1838801
commit 70c61aa53e
6 changed files with 158 additions and 16 deletions

View File

@ -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 {})

View File

@ -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):

View File

@ -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(

View File

@ -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的时候不应该裁剪

View File

@ -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"

View File

@ -58,6 +58,7 @@
### 修复
- 调整智能体 todo 展示语义:待办状态不再作为 `capabilities` 前端开关,而是直接根据运行态 `agent_state.todos` 渲染;同时将 todo 入口从 Agent Panel 移到输入框内的轻量浮层,并让右侧“状态工作台”收敛为文件系统视图,输入框按钮文案同步由“状态”调整为“文件”
- 优化 Agent 输入框 mention 行为:在保留附件 mention 的同时,将共享 `workspace` 文件纳入候选范围;并将 `@` 空查询时的候选列表改为空,仅在继续输入后再执行筛选,避免工作区文件过多时直接铺满下拉面板
- 为前端工作台文件树补齐文件删除能力:`/api/viewer/filesystem/file` 新增删除接口,`AgentPanel` 文件节点新增删除按钮与确认交互,删除后会同步刷新树与预览状态
- 调整前端工作台文件预览交互:恢复默认侧边/弹窗预览,并新增显式“全屏预览”入口;全屏模式下由预览内容直接覆盖整页,仅保留右上角悬浮关闭按钮;同时修复 HTML 文件首次在弹窗中预览偶现白屏的问题,改为在内容更新后强制重建 `iframe`