wx-agent/tools/text2sql_tool.py
2026-07-15 16:30:58 +08:00

778 lines
26 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 工具
核心思想Agent负责推理Tool负责SQLReviewer负责纠错
内部流程:
1. Schema Injection - 注入数据库表结构
2. SQL Generation - LLM生成SQL使用 SQLAlchemy 连接池 + extra_info 辅助)
3. SQL Review - LLM审查SQL正确性
4. SQL Guard - 代码级安全校验(硬限制)
5. SQL Execute - 执行SQL并返回结果使用 SQLAlchemy 连接池单例)
6. Merge - 相似结果合并Embedding
"""
import re
import json
import asyncio
from datetime import datetime
from typing import Dict, Any, Optional, List, Tuple
from sqlalchemy import text
from langchain_community.utilities import SQLDatabase
from config import POSTGRES_CONNECTION_STRING,POSTGRES_SQLALCHEMY_URL
from modelsAPI.model_api import OpenaiAPI
from prompts import TEXT2SQL_PROMPTS
from tools.texttosql_utils import get_extra_info
# ============================================================
# 数据库连接池单例(避免重复初始化)
# ============================================================
_db_instance: Optional[SQLDatabase] = None
_db_lock = asyncio.Lock()
async def _get_db() -> SQLDatabase:
"""懒加载并缓存数据库连接,避免重复初始化"""
global _db_instance
if _db_instance is None:
async with _db_lock:
if _db_instance is None:
_db_instance = await asyncio.to_thread(
SQLDatabase.from_uri, POSTGRES_SQLALCHEMY_URL
)
return _db_instance
# ============================================================
# Step 1: Schema Injection
# ============================================================
FAULT_RECORDS_SCHEMA = {
"table_name": "fault_records",
"description": "故障维修记录表,存储所有故障诊断和维修的记录",
"columns": {
"id": {
"type": "SERIAL",
"description": "自增主键",
"usable_in_query": False
},
"ship_number": {
"type": "TEXT",
"description": "舰船舷号,如'101''163''554'2-4位数字",
"usable_in_query": True,
"examples": ["101", "163", "554", "602"]
},
"device_name": {
"type": "TEXT",
"description": "设备名称,如'主机''发电机''分油机''空压机'",
"usable_in_query": True,
"examples": ["主机", "发电机", "分油机", "空压机", "舵机", "锅炉"]
},
"fault": {
"type": "TEXT",
"description": "故障现象描述,如'滑油压力低''异常振动''启动失败'",
"usable_in_query": True,
"examples": ["滑油压力低", "异常振动", "启动失败", "排烟温度高"]
},
"spare_parts": {
"type": "JSONB",
"description": "消耗的备件列表JSON数组格式'[\"密封圈\", \"滤芯\"]'",
"usable_in_query": True,
"json_query_hint": "使用 jsonb_array_elements_text(spare_parts) 展开备件进行统计"
},
"system_name": {
"type": "TEXT",
"description": "所属系统名称,如'动力系统''电气系统''导航系统'",
"usable_in_query": True,
"examples": ["动力系统", "电气系统", "导航系统", "通信系统", "其他"]
},
"created_at": {
"type": "TIMESTAMP",
"description": "记录创建时间,默认为当前时间",
"usable_in_query": True,
"examples": ["2025-01-15 10:30:00"]
}
}
}
# 字段白名单SQL Guard 使用)
ALLOWED_COLUMNS = set(FAULT_RECORDS_SCHEMA["columns"].keys())
def get_schema_prompt_section() -> str:
"""生成 Schema 注入的提示词部分"""
lines = [f"数据库表:{FAULT_RECORDS_SCHEMA['table_name']}"]
lines.append(f"说明:{FAULT_RECORDS_SCHEMA['description']}")
lines.append("")
lines.append("字段列表:")
for col_name, col_info in FAULT_RECORDS_SCHEMA["columns"].items():
if not col_info.get("usable_in_query", True):
continue
line = f" - {col_name} ({col_info['type']}): {col_info['description']}"
if "examples" in col_info:
line += f"(示例:{', '.join(col_info['examples'][:3])}"
if "json_query_hint" in col_info:
line += f"\n 查询提示:{col_info['json_query_hint']}"
lines.append(line)
return "\n".join(lines)
# ============================================================
# Step 2: SQL Generation
# ============================================================
async def generate_sql(question: str, current_date: str) -> Dict[str, Any]:
"""
使用 LLM 根据用户问题生成 SQL。
并发预热数据库连接与获取 extra_info提升整体响应速度。
Returns:
{"sql": "SELECT ...", "success": True/False, "error": "..."}
"""
schema_section = get_schema_prompt_section()
# 并发:获取辅助信息 + 预热DB连接
extra_info, _ = await asyncio.gather(
get_extra_info(question),
_get_db(),
)
prompt = TEXT2SQL_PROMPTS["sql_generation"].format(
schema=schema_section,
question=question,
current_date=current_date,
extra_info=extra_info or ""
)
print("输出最终sql的提示词内容")
print(prompt)
try:
result = await OpenaiAPI.open_api_chat_without_thinking(
prompt,
model=None,
json_output=True
)
# 提取 JSON
json_match = re.search(r'\{.*\}', result, re.DOTALL)
if json_match:
parsed = json.loads(json_match.group(0))
sql = parsed.get("sql", "").strip()
if sql:
return {"sql": sql, "success": True, "error": ""}
# 兜底:直接提取 SELECT 语句
sql_match = re.search(r'(SELECT\s+.+?;?)$', result.strip(), re.IGNORECASE | re.DOTALL)
if sql_match:
sql = sql_match.group(1).strip().rstrip(';')
return {"sql": sql, "success": True, "error": ""}
return {"sql": "", "success": False, "error": "LLM未返回有效SQL"}
except Exception as e:
return {"sql": "", "success": False, "error": f"SQL生成失败: {str(e)}"}
# ============================================================
# Step 3: SQL Review (LLM 审查)
# ============================================================
async def review_sql(question: str, sql: str) -> Dict[str, Any]:
"""
使用 LLM 审查 SQL 是否真正回答了用户问题
Returns:
{
"approved": True/False,
"issues": ["问题1", "问题2"],
"suggested_sql": "修正后的SQL如果有问题",
"review_comment": "审查意见"
}
"""
from modelsAPI.model_api import OpenaiAPI
from prompts import TEXT2SQL_PROMPTS
schema_section = get_schema_prompt_section()
prompt = TEXT2SQL_PROMPTS["sql_review"].format(
schema=schema_section,
question=question,
sql=sql
)
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))
approved = parsed.get("approved", False)
issues = parsed.get("issues", [])
suggested_sql = parsed.get("suggested_sql", "").strip()
review_comment = parsed.get("review_comment", "")
if suggested_sql:
suggested_sql = suggested_sql.rstrip(';')
return {
"approved": approved,
"issues": issues,
"suggested_sql": suggested_sql,
"review_comment": review_comment
}
return {
"approved": True,
"issues": [],
"suggested_sql": "",
"review_comment": "审查结果解析失败,默认通过"
}
except Exception as e:
return {
"approved": True,
"issues": [],
"suggested_sql": "",
"review_comment": f"审查过程异常,默认通过: {str(e)}"
}
# ============================================================
# Step 4: SQL Guard (代码级安全校验)
# ============================================================
class SQLGuardError(Exception):
"""SQL Guard 校验失败异常"""
pass
def guard_sql(sql: str) -> Tuple[bool, str, str]:
"""
SQL 安全守卫 - 代码硬限制
Returns:
(passed, cleaned_sql, error_message)
"""
if not sql or not sql.strip():
return False, "", "SQL为空"
sql_upper = sql.strip().upper()
# 1. 只允许 SELECT
if not sql_upper.startswith("SELECT"):
return False, "", "只允许SELECT查询禁止INSERT/UPDATE/DELETE/CREATE等操作"
# 2. 禁止 UNION
if "UNION" in sql_upper:
return False, "", "禁止使用UNION操作"
# 3. 禁止子查询嵌套过深最多2层
subquery_count = sql_upper.count("SELECT") - 1
if subquery_count > 2:
return False, "", f"子查询嵌套过深({subquery_count}最多允许2层"
# 4. 禁止危险函数
dangerous_functions = ["PG_SLEEP", "SLEEP", "BENCHMARK", "WAITFOR", "DELAY"]
for func in dangerous_functions:
if func in sql_upper:
return False, "", f"禁止使用危险函数: {func}"
# 5. 禁止系统表访问
system_tables = ["PG_", "INFORMATION_SCHEMA", "PG_CATALOG", "PG_CLASS"]
for sys_table in system_tables:
if sys_table in sql_upper:
return False, "", f"禁止访问系统表: {sys_table}"
# 6. 禁止注释注入
if "--" in sql or "/*" in sql or "*/" in sql:
return False, "", "SQL中不允许包含注释"
# 7. 禁止分号(防止多语句注入)
if ";" in sql.rstrip(';'):
return False, "", "禁止多语句执行"
# 8. 检查字段白名单
used_columns = set()
dot_matches = re.findall(r'fault_records\.(\w+)', sql, re.IGNORECASE)
used_columns.update(m.lower() for m in dot_matches)
select_match = re.search(r'SELECT\s+(.*?)\s+FROM', sql, re.IGNORECASE | re.DOTALL)
if select_match:
select_part = select_match.group(1)
if '*' not in select_part:
for part in select_part.split(','):
part = part.strip()
col_match = re.search(r'(\w+)\s*(?:AS|as)?\s*$', part)
if col_match:
col_name = col_match.group(1).lower()
if col_name not in ('count', 'sum', 'avg', 'max', 'min', 'as',
'desc', 'asc', 'null', 'distinct', 'all',
'jsonb_array_elements_text', 'jsonb_array_elements',
'coalesce', 'cast', 'date', 'timestamp', 'text'):
used_columns.add(col_name)
for col in used_columns:
if col not in ALLOWED_COLUMNS and col != 'id':
return False, "", f"使用了不允许的字段: {col}"
# 9. 自动添加 LIMIT如果没有
cleaned_sql = sql.strip().rstrip(';')
if "LIMIT" not in sql_upper:
cleaned_sql = f"{cleaned_sql} LIMIT 100"
# 10. 强制只查 fault_records 表
from_match = re.search(r'FROM\s+(\w+)', sql, re.IGNORECASE)
if from_match:
table_name = from_match.group(1).lower()
if table_name != "fault_records":
return False, "", f"只允许查询 fault_records 表,不允许查询: {table_name}"
return True, cleaned_sql, ""
# ============================================================
# Step 5: SQL Execute
# ============================================================
async def execute_sql(sql: str) -> Dict[str, Any]:
"""
执行 SQL 并返回结果。
使用模块级 SQLAlchemy 连接池单例,避免每次重复建立连接。
Returns:
{"rows": [...], "columns": [...], "row_count": int, "success": True/False, "error": ""}
"""
try:
db = await _get_db()
def _run() -> Dict[str, Any]:
try:
with db._engine.connect() as conn:
result = conn.execute(text(sql))
columns = list(result.keys())
result_rows = []
for row in result.fetchall():
row_dict = {}
for i, col in enumerate(columns):
val = row[i]
if isinstance(val, datetime):
val = val.isoformat()
elif isinstance(val, (list, dict)):
val = json.dumps(val, ensure_ascii=False)
elif not isinstance(val, (str, int, float, bool, type(None))):
val = str(val)
row_dict[col] = val
result_rows.append(row_dict)
conn.commit()
return {
"rows": result_rows,
"columns": columns,
"row_count": len(result_rows),
"success": True,
"error": ""
}
except Exception as e:
return {
"rows": [],
"columns": [],
"row_count": 0,
"success": False,
"error": f"SQL执行失败: {str(e)}"
}
return await asyncio.to_thread(_run)
except Exception as e:
return {
"rows": [],
"columns": [],
"row_count": 0,
"success": False,
"error": f"获取数据库连接失败: {str(e)}"
}
# ============================================================
# Text2SQL 主入口 - 内部闭环
# ============================================================
MAX_RETRY = 2 # 最大重试次数(生成→审查→修正)
async def text2sql_tool(question: str, current_date: Optional[str] = None) -> Dict[str, Any]:
"""
受控 Text2SQL 工具 - 内部形成闭环
流程question → SQL生成 → SQL审查 → SQL Guard → SQL执行 → 返回结果
Args:
question: 用户问题
current_date: 当前日期格式YYYY-MM-DD默认为今天
Returns:
{
"success": True/False,
"sql": "最终执行的SQL",
"rows": [...],
"columns": [...],
"row_count": int,
"error": "",
"review_log": [...] # 审查日志
}
"""
if not current_date:
current_date = datetime.now().strftime("%Y-%m-%d")
review_log = []
current_sql = ""
for attempt in range(MAX_RETRY + 1):
# Step 2: SQL Generation
if attempt == 0:
gen_result = await generate_sql(question, current_date)
else:
gen_result = await generate_sql_with_feedback(
question, current_date, current_sql, review_log
)
if not gen_result["success"]:
return {
"success": False,
"sql": "",
"rows": [],
"columns": [],
"row_count": 0,
"error": gen_result["error"],
"review_log": review_log
}
current_sql = gen_result["sql"]
# Step 3: SQL Review
review_result = await review_sql(question, current_sql)
review_log.append({
"attempt": attempt + 1,
"sql": current_sql,
"approved": review_result["approved"],
"issues": review_result["issues"],
"comment": review_result["review_comment"]
})
if not review_result["approved"]:
if review_result["suggested_sql"]:
current_sql = review_result["suggested_sql"]
review_log.append({
"attempt": attempt + 1,
"sql": current_sql,
"note": "使用审查建议的修正SQL"
})
elif attempt < MAX_RETRY:
continue
else:
return {
"success": False,
"sql": current_sql,
"rows": [],
"columns": [],
"row_count": 0,
"error": f"SQL审查不通过: {'; '.join(review_result['issues'])}",
"review_log": review_log
}
# Step 4: SQL Guard
passed, cleaned_sql, guard_error = guard_sql(current_sql)
if not passed:
review_log.append({
"attempt": attempt + 1,
"sql": current_sql,
"guard_passed": False,
"guard_error": guard_error
})
if attempt < MAX_RETRY:
review_log.append({"note": f"Guard拒绝: {guard_error},将重新生成"})
continue
else:
return {
"success": False,
"sql": current_sql,
"rows": [],
"columns": [],
"row_count": 0,
"error": f"SQL安全校验失败: {guard_error}",
"review_log": review_log
}
current_sql = cleaned_sql
review_log.append({
"attempt": attempt + 1,
"sql": current_sql,
"guard_passed": True,
"note": "Guard通过"
})
# Step 5: SQL Execute
exec_result = await execute_sql(current_sql)
if not exec_result["success"]:
if attempt < MAX_RETRY:
review_log.append({
"attempt": attempt + 1,
"sql": current_sql,
"exec_error": exec_result["error"],
"note": "执行失败,将重新生成"
})
continue
else:
return {
"success": False,
"sql": current_sql,
"rows": [],
"columns": [],
"row_count": 0,
"error": exec_result["error"],
"review_log": review_log
}
# Step 6: 相似结果合并
merged_rows = await merge_similar_rows(
exec_result["rows"], exec_result["columns"], question
)
return {
"success": True,
"sql": current_sql,
"rows": merged_rows,
"columns": exec_result["columns"],
"row_count": len(merged_rows),
"raw_row_count": exec_result["row_count"],
"error": "",
"review_log": review_log
}
return {
"success": False,
"sql": current_sql,
"rows": [],
"columns": [],
"row_count": 0,
"error": "超过最大重试次数",
"review_log": review_log
}
async def generate_sql_with_feedback(
question: str,
current_date: str,
previous_sql: str,
review_log: List[Dict[str, Any]]
) -> Dict[str, Any]:
"""
基于审查反馈重新生成 SQL
"""
from modelsAPI.model_api import OpenaiAPI
from prompts import TEXT2SQL_PROMPTS
schema_section = get_schema_prompt_section()
feedback_parts = []
for log_entry in review_log:
if "issues" in log_entry and log_entry["issues"]:
feedback_parts.append(f"问题: {'; '.join(log_entry['issues'])}")
if "guard_error" in log_entry:
feedback_parts.append(f"安全校验失败: {log_entry['guard_error']}")
if "exec_error" in log_entry:
feedback_parts.append(f"执行失败: {log_entry['exec_error']}")
feedback = "\n".join(feedback_parts) if feedback_parts else "无具体反馈"
prompt = TEXT2SQL_PROMPTS["sql_regeneration"].format(
schema=schema_section,
question=question,
current_date=current_date,
previous_sql=previous_sql,
feedback=feedback
)
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))
sql = parsed.get("sql", "").strip()
if sql:
return {"sql": sql, "success": True, "error": ""}
sql_match = re.search(r'(SELECT\s+.+?;?)$', result.strip(), re.IGNORECASE | re.DOTALL)
if sql_match:
sql = sql_match.group(1).strip().rstrip(';')
return {"sql": sql, "success": True, "error": ""}
return {"sql": "", "success": False, "error": "重新生成SQL失败"}
except Exception as e:
return {"sql": "", "success": False, "error": f"重新生成SQL异常: {str(e)}"}
# ============================================================
# Step 6: 相似结果合并Embedding
# ============================================================
MERGE_KEYWORDS = ["故障", "频次", "排名", "高频", "最多", "top", "发生次数"]
def _should_merge(columns: List[str], question: str) -> bool:
"""判断是否需要合并相似行"""
question_lower = question.lower()
if any(kw in question_lower for kw in MERGE_KEYWORDS):
col_names = [c.lower() for c in columns]
has_name = any(n in col_names for n in ["name", "故障", "fault", "device_name"])
has_count = any(n in col_names for n in ["count", "次数", "频次"])
return has_name and has_count
return False
def _find_name_column(columns: List[str]) -> Optional[str]:
"""找到名称列"""
priority = ["name", "故障名称", "故障", "fault", "device_name", "备件", "spare"]
col_lower_map = {c.lower(): c for c in columns}
for p in priority:
if p in col_lower_map:
return col_lower_map[p]
for c in columns:
if c.lower() not in ("count", "次数", "频次", "rank", "排名", "system_name", "系统"):
return c
return None
def _find_count_column(columns: List[str]) -> Optional[str]:
"""找到计数列"""
priority = ["count", "次数", "频次", "cnt"]
col_lower_map = {c.lower(): c for c in columns}
for p in priority:
if p in col_lower_map:
return col_lower_map[p]
return None
def _find_system_column(columns: List[str]) -> Optional[str]:
"""找到系统分类列"""
priority = ["system_name", "系统", "系统分类", "system"]
col_lower_map = {c.lower(): c for c in columns}
for p in priority:
if p in col_lower_map:
return col_lower_map[p]
return None
async def merge_similar_rows(
rows: List[Dict[str, Any]],
columns: List[str],
question: str,
similarity_threshold: float = 0.85
) -> List[Dict[str, Any]]:
"""
使用 embedding 合并相似的故障/备件行
例如:
- "连杆大端轴承发生损坏异常" (7次, 其他) + "连杆大端轴承发生损坏异常" (6次, 汽车配件)
"连杆大端轴承发生损坏异常" (13次, 其他)
- "发动机中空冷器管束芯体结垢或堵塞..." + "空冷器管束芯体结垢或堵塞..."
→ 合并为一条
"""
if not rows or len(rows) <= 1:
return rows
if not _should_merge(columns, question):
return rows
name_col = _find_name_column(columns)
count_col = _find_count_column(columns)
system_col = _find_system_column(columns)
if not name_col or not count_col:
return rows
try:
from modelsAPI.model_api import OpenaiAPI
names = [row[name_col] for row in rows if row.get(name_col)]
if len(names) <= 1:
return rows
embeddings = await OpenaiAPI.get_embeddings_batch_async(names)
name_emb_map = {name: emb for name, emb in zip(names, embeddings)}
sorted_rows = sorted(rows, key=lambda r: int(r.get(count_col, 0)), reverse=True)
merged_rows = []
merged_indices = set()
for i, row in enumerate(sorted_rows):
if i in merged_indices:
continue
name_i = row.get(name_col, "")
emb_i = name_emb_map.get(name_i)
if not emb_i:
merged_rows.append(row)
continue
merged_row = dict(row)
total_count = int(row.get(count_col, 0))
systems = set()
if system_col and row.get(system_col):
systems.add(row[system_col])
for j in range(i + 1, len(sorted_rows)):
if j in merged_indices:
continue
name_j = sorted_rows[j].get(name_col, "")
emb_j = name_emb_map.get(name_j)
if not emb_j:
continue
similarity = OpenaiAPI.cosine_similarity(emb_i, emb_j)
if similarity >= similarity_threshold:
total_count += int(sorted_rows[j].get(count_col, 0))
if system_col and sorted_rows[j].get(system_col):
systems.add(sorted_rows[j][system_col])
merged_indices.add(j)
merged_row[count_col] = total_count
if system_col and systems:
merged_row[system_col] = "".join(sorted(systems))
merged_rows.append(merged_row)
merged_indices.add(i)
merged_rows.sort(key=lambda r: int(r.get(count_col, 0)), reverse=True)
if "rank" in columns or "排名" in columns:
rank_col = "rank" if "rank" in columns else "排名"
for idx, row in enumerate(merged_rows):
row[rank_col] = idx + 1
return merged_rows
except Exception as e:
print(f"Embedding合并失败使用原始数据: {str(e)}")
return rows