""" 工作流:受控 ReAct Text2SQL 统计智能体 架构(简洁直接,无多轮交互): 用户问题 ↓ Router Node(判断:统计分析 / 维修反馈 / 普通问答) ├── 普通问答 → final_answer → END ├── 维修反馈 → repair_feedback → final_answer → END └── 统计分析 → text2sql_tool(内部闭环)→ final_answer → END text2sql_tool 内部闭环:SQL生成 → SQL审查 → SQL Guard → SQL执行 """ import sys import os import json import re from datetime import datetime from typing import TypedDict, Optional, Dict, Any, List sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from langgraph.graph import StateGraph, START, END from modelsAPI.model_api import OpenaiAPI from prompts import TEXT2SQL_PROMPTS from tools.text2sql_tool import text2sql_tool from tools.fault_statistics import repair_feedback_statistics_tool # ============================================================ # 状态定义 # ============================================================ class StatisticsState(TypedDict): # 输入 extracted_text: str combined_query: str route_flag: str history_message: str history_messages: List[Dict[str, Any]] # 路由结果 route_result: Optional[str] # "statistics" / "feedback" / "normal" # Text2SQL 结果 sql: Optional[str] sql_result: Optional[Dict[str, Any]] # 维修反馈结果 feedback_result: Optional[Dict[str, Any]] feedback_params: Optional[Dict[str, Any]] # 输出 response: str actions: List[Dict[str, Any]] suggestedReplies: List[Dict[str, Any]] error_message: str # ============================================================ # Node 1: Router Node # ============================================================ async def router_node(state: StatisticsState) -> Dict[str, Any]: """ 路由判断:统计分析 / 维修反馈 / 普通问答 """ question = state.get("combined_query") or state.get("extracted_text", "") prompt = TEXT2SQL_PROMPTS["router"].format(question=question) try: result = await OpenaiAPI.open_api_chat_without_thinking( prompt, model=None, json_output=True ) json_match = re.search(r'\{.*\}', result, re.DOTALL) if json_match: parsed = json.loads(json_match.group(0)) route = parsed.get("route", "statistics").strip().lower() if route in ("statistics", "feedback", "normal"): return {"route_result": route} except Exception as e: print(f"路由判断失败: {str(e)}") return {"route_result": "statistics"} # ============================================================ # Node 2: Text2SQL Node # ============================================================ async def text2sql_node(state: StatisticsState) -> Dict[str, Any]: """ 调用受控 Text2SQL 工具 内部闭环:SQL生成 → SQL审查 → SQL Guard → SQL执行(含重试) """ question = state.get("combined_query") or state.get("extracted_text", "") current_date = datetime.now().strftime("%Y-%m-%d") result = await text2sql_tool(question=question, current_date=current_date) return { "sql": result.get("sql", ""), "sql_result": result } # ============================================================ # Node 3: Repair Feedback Node # ============================================================ async def repair_feedback_node(state: StatisticsState) -> Dict[str, Any]: """ 处理维修反馈类问题(走外部API) """ question = state.get("combined_query") or state.get("extracted_text", "") params = await _extract_feedback_params(question) try: result = await repair_feedback_statistics_tool.ainvoke(params) return {"feedback_result": result, "feedback_params": params} except Exception as e: return {"feedback_result": {"success": False, "error": str(e)}, "feedback_params": params} async def _extract_feedback_params(question: str) -> Dict[str, Any]: """从用户问题中提取维修反馈参数""" current_date = datetime.now().strftime("%Y-%m-%d") prompt = f"""从以下问题中提取维修反馈统计参数,返回JSON格式。 当前日期:{current_date} 用户问题:{question} 需要提取的参数: - source: 来源筛选(好评/差评/反馈),如果没有则为空字符串 - start_date: 开始日期(格式:YYYY-MM-DD),如果没有则为空字符串 - end_date: 结束日期(格式:YYYY-MM-DD),如果没有则为空字符串 时间映射: - "近一个月" → start_date = 当前日期前30天 - "近一周" → start_date = 当前日期前7天 - "近三个月" → start_date = 当前日期前90天 输出格式: {{"source": "", "start_date": "", "end_date": ""}} 只输出JSON。""" try: result = await OpenaiAPI.open_api_chat_without_thinking( prompt, model=None, json_output=True ) json_match = re.search(r'\{.*\}', result, re.DOTALL) if json_match: parsed = json.loads(json_match.group(0)) params = {} if parsed.get("source"): params["source"] = parsed["source"] if parsed.get("start_date"): params["start_date"] = parsed["start_date"] if parsed.get("end_date"): params["end_date"] = parsed["end_date"] return params except Exception: pass return {} # ============================================================ # Node 4: Final Answer Node # ============================================================ async def final_answer_node(state: StatisticsState) -> Dict[str, Any]: """ 将查询结果转为自然语言回答 """ route_result = state.get("route_result", "statistics") question = state.get("combined_query") or state.get("extracted_text", "") if route_result == "feedback": return await _format_feedback_answer(state, question) elif route_result == "normal": return { "response": "您的问题似乎不属于统计分析类问题,请尝试更具体的统计查询,例如:\"最近一个月哪个故障最多\"、\"101舰备件消耗情况\"等。", "actions": [], "suggestedReplies": [ {"title": "故障频次统计", "content": "最近一个月哪个故障最多"}, {"title": "备件消耗统计", "content": "101舰近一个月备件消耗"}, {"title": "维修反馈", "content": "近一个月维修差评情况"} ] } else: return await _format_statistics_answer(state, question) async def _format_statistics_answer(state: StatisticsState, question: str) -> Dict[str, Any]: """格式化统计分析结果""" sql_result = state.get("sql_result", {}) if not sql_result.get("success", False): error = sql_result.get("error", "未知错误") return { "response": f"统计分析执行失败:{error}", "actions": [], "suggestedReplies": [ {"title": "故障频次统计", "content": "最近一个月哪个故障最多"}, {"title": "备件消耗统计", "content": "近一个月备件消耗排名"} ] } rows = sql_result.get("rows", []) sql = sql_result.get("sql", "") row_count = sql_result.get("row_count", 0) # 使用 LLM 生成自然语言回答 result_data = json.dumps(rows[:20], ensure_ascii=False, indent=2) prompt = TEXT2SQL_PROMPTS["final_answer"].format( question=question, sql=sql, row_count=row_count, result_data=result_data ) try: answer = await OpenaiAPI.open_api_chat_without_thinking(prompt, model=None) response = answer.strip() except Exception: response = _simple_format_result(rows, row_count) return { "response": response, "actions": [], "suggestedReplies": _generate_suggested_replies(question) } async def _format_feedback_answer(state: StatisticsState, question: str) -> Dict[str, Any]: """格式化维修反馈结果""" feedback_result = state.get("feedback_result", {}) feedback_params = state.get("feedback_params", {}) if not feedback_result.get("success", False): return { "response": "维修反馈统计失败,请稍后重试。", "actions": [], "suggestedReplies": [] } feedback_data = feedback_result.get("feedbackData", []) total_count = feedback_result.get("totalCount", 0) source_distribution = feedback_result.get("sourceDistribution", {}) source = feedback_params.get("source", "") start_date = feedback_params.get("start_date", "") end_date = feedback_params.get("end_date", "") if not feedback_data: source_desc = f"【{source}】" if source else "" time_desc = f"({start_date} 至 {end_date})" if start_date and end_date else "" return { "response": f"未找到{source_desc}维修反馈统计数据{time_desc}。", "actions": [], "suggestedReplies": [ {"title": "查看差评", "content": "统计近一个月维修差评情况"}, {"title": "查看好评", "content": "统计近一个月维修好评情况"} ] } response_parts = [] source_desc = f"【{source}】" if source else "" time_desc = f"({start_date} 至 {end_date})" if start_date and end_date else "" response_parts.append(f"为您统计{source_desc}的维修反馈数据{time_desc}。\n") if source_distribution: dist_str = "、".join([f"{k} {v} 条" for k, v in source_distribution.items()]) response_parts.append(f"核心结论:共收到维修反馈 {total_count} 条({dist_str})。\n") response_parts.append("排名 | 标题 | 来源 | 反馈人 | 时间") response_parts.append("--- | --- | --- | --- | ---") for item in feedback_data: rank = item.get("rank", 0) title = str(item.get("title", ""))[:30] item_source = str(item.get("source", "未知")) user_name = str(item.get("user_name", "-")) created_at = str(item.get("created_at", "-"))[:10] response_parts.append(f"{rank} | {title} | {item_source} | {user_name} | {created_at}") return { "response": "\n".join(response_parts), "actions": [], "suggestedReplies": [ {"title": "查看差评", "content": "统计近一个月维修差评情况"}, {"title": "查看好评", "content": "统计近一个月维修好评情况"}, {"title": "查看近一周反馈", "content": "统计近一周维修反馈情况"} ] } def _simple_format_result(rows: List[Dict[str, Any]], row_count: int) -> str: """简单格式化查询结果(LLM生成失败时的兜底)""" if not rows: return "未查询到相关数据。" columns = list(rows[0].keys()) parts = [f"共查询到 {row_count} 条记录。\n"] parts.append(" | ".join(columns)) parts.append(" | ".join(["---"] * len(columns))) for row in rows[:10]: values = [str(row.get(col, ""))[:30] for col in columns] parts.append(" | ".join(values)) if row_count > 10: parts.append(f"\n... 还有 {row_count - 10} 条记录") return "\n".join(parts) def _generate_suggested_replies(question: str) -> List[Dict[str, Any]]: """根据问题生成建议回复""" return [ {"title": "故障频次统计", "content": "最近一个月哪个故障最多"}, {"title": "备件消耗统计", "content": "近一个月备件消耗排名"}, {"title": "按系统统计", "content": "动力系统近一个月故障统计"}, {"title": "按舷号统计", "content": "101舰近一个月故障统计"} ] # ============================================================ # 条件路由 # ============================================================ def route_after_router(state: StatisticsState) -> str: """Router 之后的路由""" route = state.get("route_result", "statistics") if route == "feedback": return "feedback" elif route == "normal": return "normal" else: return "statistics" # ============================================================ # 构建工作流 # ============================================================ def create_statistics_workflow(): """ 创建受控 Text2SQL 统计工作流 架构(简洁直线型,无多轮交互): START ↓ router_node ├── normal → final_answer_node → END ├── feedback → repair_feedback_node → final_answer_node → END └── statistics → text2sql_node → final_answer_node → END """ workflow = StateGraph(StatisticsState) # 添加节点 workflow.add_node("router_node", router_node) workflow.add_node("text2sql_node", text2sql_node) workflow.add_node("repair_feedback_node", repair_feedback_node) workflow.add_node("final_answer_node", final_answer_node) # 入口 workflow.add_edge(START, "router_node") # Router 分支 workflow.add_conditional_edges( "router_node", route_after_router, { "statistics": "text2sql_node", "feedback": "repair_feedback_node", "normal": "final_answer_node" } ) # 各分支 → final_answer → END workflow.add_edge("text2sql_node", "final_answer_node") workflow.add_edge("repair_feedback_node", "final_answer_node") workflow.add_edge("final_answer_node", END) return workflow.compile() # ============================================================ # 统一入口 # ============================================================ async def run_statistics_workflow( extracted_text: str, combined_query: str, history_message: str = "", route_flag: str = "" ) -> Dict[str, Any]: """ 执行统计工作流(统一接口) """ app = create_statistics_workflow() initial_state = { "extracted_text": extracted_text, "combined_query": combined_query, "route_flag": route_flag, "history_message": history_message, "history_messages": [], "route_result": None, "sql": None, "sql_result": None, "feedback_result": None, "feedback_params": None, "response": "", "actions": [], "suggestedReplies": [], "error_message": "" } try: final_state = await app.ainvoke(initial_state) if final_state.get("error_message"): return { "response": f"统计工作流执行失败: {final_state['error_message']}", "actions": [], "result_tag": "statistics", "suggestedReplies": [] } return { "response": final_state.get("response", ""), "actions": final_state.get("actions", []), "result_tag": "statistics", "suggestedReplies": final_state.get("suggestedReplies", []) } except Exception as e: return { "response": f"统计工作流执行失败: {str(e)}", "actions": [], "result_tag": "statistics", "suggestedReplies": [] } # ============================================================ # 测试 # ============================================================ if __name__ == "__main__": import asyncio test_cases = [ "最近一个月哪个故障最多", "统计101舰的备件消耗情况", "近三个月动力系统故障频次排名", "统计维修反馈差评情况", "连杆异常磨损怎么修", ] async def test(): for i, test_text in enumerate(test_cases, 1): print(f"\n{'='*60}") print(f"测试用例 {i}: {test_text}") print(f"{'='*60}") result = await run_statistics_workflow( extracted_text=test_text, combined_query=test_text ) print(f"\n响应:") print(result.get("response", "无响应")) asyncio.run(test())