QueryTool.java
9.04 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
package com.xly.tool;
import com.xly.service.AuditService;
import dev.langchain4j.agent.tool.P;
import dev.langchain4j.agent.tool.Tool;
import dev.langchain4j.model.ollama.OllamaChatModel;
import net.sf.jsqlparser.parser.CCJSqlParserUtil;
import net.sf.jsqlparser.statement.Statement;
import net.sf.jsqlparser.statement.select.Select;
import org.springframework.beans.factory.annotation.Qualifier;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.stereotype.Component;
import java.util.List;
import java.util.Map;
/**
* Query 工具:**只读 SQL 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题
* (跨表汇总、按条件计数排名等)。
*
* <p>安全栈:用 coder 模型据 KG 字段字典接地生成 SQL → jsqlparser 强制**单条 SELECT** →
* 挡 {@code INTO OUTFILE / LOAD_FILE / information_schema / SLEEP / BENCHMARK} 与多语句 →
* 强制 LIMIT。(本地单品牌,租户注入留作生产加固;见架构 §9。)SQL 入审计。
*/
@Component
public class QueryTool {
private final OllamaChatModel sqlModel;
private final JdbcTemplate jdbc;
private final AuditService audit;
public QueryTool(@Qualifier("sqlChatModel") OllamaChatModel sqlModel, JdbcTemplate jdbc, AuditService audit) {
this.sqlModel = sqlModel;
this.jdbc = jdbc;
this.audit = audit;
}
@Tool("用**只读 SQL** 回答没有现成表单/记录能直接答的临时统计或分析问题"
+ "(如跨表汇总、按条件计数、排名、分组统计)。仅在 readFormData / lookupRecord 无法回答时才用。")
public String queryData(@P("用自然语言描述要统计/分析什么") String question) {
if (question == null || question.isBlank()) {
return "请描述要查询统计的内容。";
}
String hint = schemaHint(question);
String sql = null;
String lastErr = null;
// 自修复重试:SQL 校验/执行报错就把错误喂回模型重新生成,最多 3 次
for (int attempt = 0; attempt < 3; attempt++) {
try {
sql = cleanSql(sqlModel.chat(buildPrompt(hint, question, sql, lastErr)));
} catch (Exception e) {
return "生成查询失败:" + e.getMessage();
}
String reject = validate(sql);
if (reject != null) {
lastErr = reject;
if (attempt < 2) {
continue;
}
audit.log(null, null, "query", "REJECTED", sql, false, reject);
return "无法安全执行该查询(" + reject + ")。可以换个更具体的问法。";
}
String limited = forceLimit(sql);
try {
List<Map<String, Object>> rows = jdbc.queryForList(limited);
audit.log(null, null, "query", "ok", limited, true, "rows=" + rows.size() + (attempt > 0 ? " (retry " + attempt + ")" : ""));
return formatRows(rows);
} catch (Exception e) {
lastErr = rootMsg(e);
if (attempt < 2) {
continue; // 下一轮把错误喂回模型自修复
}
audit.log(null, null, "query", "fail", limited, false, lastErr);
return "查询执行失败(已尝试自修复):" + lastErr;
}
}
return "查询失败。";
}
private String buildPrompt(String hint, String question, String prevSql, String prevErr) {
String repair = (prevSql == null || prevErr == null) ? "" :
"\n\n上一条 SQL:\n" + prevSql + "\n执行/校验报错:" + prevErr +
"\n请**修正该错误**后重新生成一条正确的 SELECT(注意用对表名列名、别用中文别名)。";
return """
你是 MySQL 专家。根据【问题】生成 **一条** MySQL SELECT 查询来回答它。
数据库 = xlyweberp_saas。可用的表和字段(列名=中文名):
%s
规则:只用 SELECT(严禁任何写操作 / 文件操作);需要时 JOIN;务必带合适的 LIMIT(<=100);
**列别名一律用英文**(如 cnt、total、name),ORDER BY 用英文列名或序号,**绝不要用中文做别名**;
表名、列名一律用上面给定的英文名。**只输出 SQL 本身**,不要解释、不要 markdown 代码围栏。
问题:%s%s
""".formatted(hint, question, repair);
}
private String rootMsg(Throwable e) {
Throwable r = e;
while (r.getCause() != null && r.getCause() != r) {
r = r.getCause();
}
String m = r.getMessage();
return m == null ? e.toString() : (m.length() > 300 ? m.substring(0, 300) : m);
}
/** 据字段字典把问题里出现的中文术语接地到具体表+列,喂给 coder 模型。 */
private String schemaHint(String question) {
StringBuilder sb = new StringBuilder();
try {
List<Map<String, Object>> rows = jdbc.queryForList(
"SELECT fd.sTable, " +
"GROUP_CONCAT(DISTINCT CONCAT(fd.sField,'=',fd.sChinese) ORDER BY fd.iFormUses DESC SEPARATOR ', ') cols, " +
"SUM(fd.iFormUses) usage_ " +
"FROM viw_kg_field_dict fd " +
"WHERE CHAR_LENGTH(fd.sChinese)>=2 AND INSTR(?, fd.sChinese)>0 " +
"AND fd.sTable NOT LIKE 'viw%' " +
"AND fd.sTable IN (SELECT DISTINCT sDataSource FROM viw_ai_useful_forms) " +
"GROUP BY fd.sTable ORDER BY usage_ DESC, COUNT(*) DESC LIMIT 6", question);
for (Map<String, Object> r : rows) {
String cols = String.valueOf(r.get("cols"));
if (cols.length() > 400) {
cols = cols.substring(0, 400) + "…";
}
sb.append("- ").append(r.get("sTable")).append("(").append(cols).append(")\n");
}
} catch (Exception ignore) {
}
if (sb.length() == 0) {
sb.append("(未匹配到具体表;请在问题里使用业务术语,如 客户 / 订单 / 金额 / 数量)\n");
}
return sb.toString();
}
private String cleanSql(String raw) {
if (raw == null) {
return "";
}
String s = raw.replace("```sql", "").replace("```", "").trim();
int i = s.toLowerCase().indexOf("select");
if (i > 0) {
s = s.substring(i);
}
int semi = s.indexOf(';');
if (semi >= 0) {
s = s.substring(0, semi);
}
return s.trim();
}
/** 单条 SELECT + 挡危险构造。返回 null=通过,否则=拒绝原因。 */
private String validate(String sql) {
if (sql == null || sql.isBlank()) {
return "未生成SQL";
}
String low = sql.toLowerCase();
String[] bad = {"into outfile", "into dumpfile", "load_file", "load data",
"information_schema", "sleep(", "benchmark(", "sys.", "mysql."};
for (String b : bad) {
if (low.contains(b)) {
return "含禁止构造: " + b;
}
}
try {
Statement stmt = CCJSqlParserUtil.parse(sql);
if (!(stmt instanceof Select)) {
return "只允许 SELECT";
}
} catch (Exception e) {
return "SQL 解析失败";
}
return null;
}
private String forceLimit(String sql) {
String low = sql.toLowerCase();
if (!low.matches("(?s).*\\blimit\\b.*")) {
return sql.trim() + " LIMIT 100";
}
return sql;
}
private String formatRows(List<Map<String, Object>> rows) {
if (rows.isEmpty()) {
return "查询完成,没有匹配的数据。";
}
List<String> cols = List.copyOf(rows.get(0).keySet());
StringBuilder sb = new StringBuilder();
sb.append("查询结果(").append(rows.size()).append(" 行):\n\n");
sb.append("| ").append(String.join(" | ", cols)).append(" |\n");
sb.append("|").append(" --- |".repeat(cols.size())).append("\n");
int shown = 0;
for (Map<String, Object> r : rows) {
if (shown++ >= 30) {
sb.append("| … 仅显示前 30 行 |").append(" |".repeat(Math.max(0, cols.size() - 1))).append("\n");
break;
}
StringBuilder line = new StringBuilder("| ");
for (String c : cols) {
Object v = r.get(c);
String s = v == null ? "" : v.toString().replace("|", "/").replace("\n", " ");
if (s.length() > 30) {
s = s.substring(0, 30) + "…";
}
line.append(s).append(" | ");
}
sb.append(line).append("\n");
}
return sb.toString();
}
}