wx-agent/workflows/workflow_statistics.py
2026-07-02 10:29:14 +08:00

485 lines
16 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
工作流:受控 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())