refactor: simplify chat stream helper flows

This commit is contained in:
肖泽涛 2026-02-25 16:30:36 +08:00
parent 81663e7fba
commit b7edef1042

View File

@ -3,7 +3,7 @@ import json
import traceback import traceback
import uuid import uuid
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from datetime import datetime from datetime import UTC, datetime
from langchain.messages import AIMessage, AIMessageChunk, HumanMessage from langchain.messages import AIMessage, AIMessageChunk, HumanMessage
from langgraph.types import Command from langgraph.types import Command
@ -41,7 +41,7 @@ def _build_state_files(attachments: list[dict]) -> dict:
if not file_path or not markdown: if not file_path or not markdown:
continue continue
now = datetime.utcnow().isoformat() + "+00:00" now = datetime.now(UTC).isoformat()
# 将 markdown 内容按行拆分 # 将 markdown 内容按行拆分
content_lines = markdown.split("\n") content_lines = markdown.split("\n")
files[file_path] = { files[file_path] = {
@ -194,6 +194,59 @@ async def save_messages_from_langgraph_state(
logger.error(traceback.format_exc()) logger.error(traceback.format_exc())
def _extract_interrupt_info(state) -> dict | None:
"""从 LangGraph state 中提取中断信息"""
if hasattr(state, "tasks") and state.tasks:
for task in state.tasks:
if hasattr(task, "interrupts") and task.interrupts:
return task.interrupts[0]
interrupt_data = state.values.get("__interrupt__")
if isinstance(interrupt_data, list) and interrupt_data:
return interrupt_data[0]
return None
def _get_interrupt_fields(info) -> tuple[str, str]:
"""从中断信息中提取 question 和 operation"""
defaults = ("是否批准以下操作?", "需要人工审批的操作")
if isinstance(info, dict):
return info.get("question", defaults[0]), info.get("operation", defaults[1])
return getattr(info, "question", defaults[0]), getattr(info, "operation", defaults[1])
def _ensure_full_msg(full_msg: AIMessage | None, accumulated_content: list[str]) -> AIMessage | None:
"""如果 full_msg 为空且有累积内容,构建 AIMessage"""
if not full_msg and accumulated_content:
return AIMessage(content="".join(accumulated_content))
return full_msg
async def _resolve_agent_config(
db, agent_id: str, department_id, user_id: str, agent_config_id: int | str | None
) -> tuple:
"""解析 agent_config返回 (config_item, agent_config_id)"""
config_repo = AgentConfigRepository(db)
config_item = None
if agent_config_id is not None:
try:
config_item = await config_repo.get_by_id(int(agent_config_id))
except Exception:
logger.warning(f"Failed to fetch agent config {agent_config_id}: {traceback.format_exc()}")
config_item = None
if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id):
config_item = None
if config_item is None:
config_item = await config_repo.get_or_create_default(
department_id=department_id, agent_id=agent_id, created_by=user_id
)
agent_config_id = config_item.id
return config_item, agent_config_id
async def check_and_handle_interrupts( async def check_and_handle_interrupts(
agent, agent,
langgraph_config: dict, langgraph_config: dict,
@ -208,29 +261,9 @@ async def check_and_handle_interrupts(
if not state or not state.values: if not state or not state.values:
return return
interrupt_info = None interrupt_info = _extract_interrupt_info(state)
if hasattr(state, "tasks") and state.tasks:
for task in state.tasks:
if hasattr(task, "interrupts") and task.interrupts:
interrupt_info = task.interrupts[0]
break
if not interrupt_info and state.values:
interrupt_data = state.values.get("__interrupt__")
if interrupt_data and isinstance(interrupt_data, list) and len(interrupt_data) > 0:
interrupt_info = interrupt_data[0]
if interrupt_info: if interrupt_info:
question = "是否批准以下操作?" question, operation = _get_interrupt_fields(interrupt_info)
operation = "需要人工审批的操作"
if isinstance(interrupt_info, dict):
question = interrupt_info.get("question", question)
operation = interrupt_info.get("operation", operation)
elif hasattr(interrupt_info, "question"):
question = getattr(interrupt_info, "question", question)
operation = getattr(interrupt_info, "operation", operation)
meta["interrupt"] = { meta["interrupt"] = {
"question": question, "question": question,
"operation": operation, "operation": operation,
@ -311,22 +344,9 @@ async def stream_agent_chat(
return return
agent_config_id = config.get("agent_config_id") agent_config_id = config.get("agent_config_id")
config_repo = AgentConfigRepository(db) config_item, agent_config_id = await _resolve_agent_config(
config_item = None db, agent_id, department_id, user_id, agent_config_id
if agent_config_id is not None: )
try:
config_item = await config_repo.get_by_id(int(agent_config_id))
except Exception:
logger.warning(f"Failed to fetch agent config {agent_config_id}: {traceback.format_exc()}")
config_item = None
if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id):
config_item = None
if config_item is None:
config_item = await config_repo.get_or_create_default(
department_id=department_id, agent_id=agent_id, created_by=user_id
)
agent_config_id = config_item.id
if not (thread_id := config.get("thread_id")): if not (thread_id := config.get("thread_id")):
thread_id = str(uuid.uuid4()) thread_id = str(uuid.uuid4())
@ -410,8 +430,7 @@ async def stream_agent_chat(
except Exception as e: except Exception as e:
logger.error(f"Error processing tool message: {e}") logger.error(f"Error processing tool message: {e}")
if not full_msg and accumulated_content: full_msg = _ensure_full_msg(full_msg, accumulated_content)
full_msg = AIMessage(content="".join(accumulated_content))
if conf.enable_content_guard and hasattr(full_msg, "content") and await content_guard.check(full_msg.content): if conf.enable_content_guard and hasattr(full_msg, "content") and await content_guard.check(full_msg.content):
await save_partial_message(conv_repo, thread_id, full_msg, "content_guard_blocked") await save_partial_message(conv_repo, thread_id, full_msg, "content_guard_blocked")
@ -447,8 +466,7 @@ async def stream_agent_chat(
async def save_cleanup(): async def save_cleanup():
nonlocal full_msg nonlocal full_msg
if not full_msg and accumulated_content: full_msg = _ensure_full_msg(full_msg, accumulated_content)
full_msg = AIMessage(content="".join(accumulated_content))
async with pg_manager.get_async_session_context() as new_db: async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db) new_conv_repo = ConversationRepository(new_db)
@ -476,8 +494,7 @@ async def stream_agent_chat(
error_msg = f"Error streaming messages: {e}" error_msg = f"Error streaming messages: {e}"
error_type = "unexpected_error" error_type = "unexpected_error"
if not full_msg and accumulated_content: full_msg = _ensure_full_msg(full_msg, accumulated_content)
full_msg = AIMessage(content="".join(accumulated_content))
async with pg_manager.get_async_session_context() as new_db: async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db) new_conv_repo = ConversationRepository(new_db)
@ -537,22 +554,9 @@ async def stream_agent_resume(
return return
agent_config_id = (config or {}).get("agent_config_id") agent_config_id = (config or {}).get("agent_config_id")
config_repo = AgentConfigRepository(db) config_item, agent_config_id = await _resolve_agent_config(
config_item = None db, agent_id, department_id, user_id, agent_config_id
if agent_config_id is not None: )
try:
config_item = await config_repo.get_by_id(int(agent_config_id))
except Exception:
logger.warning(f"Failed to fetch agent config {agent_config_id}: {traceback.format_exc()}")
config_item = None
if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id):
config_item = None
if config_item is None:
config_item = await config_repo.get_or_create_default(
department_id=department_id, agent_id=agent_id, created_by=user_id
)
agent_config_id = config_item.id
input_context = { input_context = {
"user_id": user_id, "user_id": user_id,