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