485 lines
16 KiB
Python
485 lines
16 KiB
Python
"""
|
||
工作流:受控 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())
|
||
|