QueryTool.java
14.9 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
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
package com.xly.tool;
import com.xly.agent.AgentIdentity;
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.expression.StringValue;
import net.sf.jsqlparser.expression.operators.conditional.AndExpression;
import net.sf.jsqlparser.expression.operators.relational.EqualsTo;
import net.sf.jsqlparser.parser.CCJSqlParserUtil;
import net.sf.jsqlparser.schema.Column;
import net.sf.jsqlparser.schema.Table;
import net.sf.jsqlparser.statement.Statement;
import net.sf.jsqlparser.statement.select.PlainSelect;
import net.sf.jsqlparser.statement.select.Select;
import net.sf.jsqlparser.util.TablesNamesFinder;
import org.springframework.jdbc.core.JdbcTemplate;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
/**
* Query 工具:**只读 SQL 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题
* (跨表汇总、按条件计数排名等)。
*
* <p>安全栈:用 coder 模型据 KG 字段字典接地生成 SQL → jsqlparser 强制**单条 SELECT** →
* 挡 {@code INTO OUTFILE / LOAD_FILE / information_schema / SLEEP / BENCHMARK} 与多语句 →
* 强制 LIMIT。(本地单品牌,租户注入留作生产加固;见架构 §9。)SQL 入审计。
*/
public class QueryTool {
private final OllamaChatModel sqlModel;
private final JdbcTemplate jdbc;
private final AuditService audit;
private final AgentIdentity identity;
public QueryTool(OllamaChatModel sqlModel, JdbcTemplate jdbc, AuditService audit, AgentIdentity identity) {
this.sqlModel = sqlModel;
this.jdbc = jdbc;
this.audit = audit;
this.identity = identity;
}
@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(applyTenant(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);**含 sBrandsId 列的表务必加 `sBrandsId='%s'` 过滤**(本企业数据);
**列别名一律用英文**(如 cnt、total、name),ORDER BY 用英文列名或序号,**绝不要用中文做别名**;
表名、列名一律用上面给定的英文名。**只输出 SQL 本身**,不要解释、不要 markdown 代码围栏。
问题:%s%s
""".formatted(hint, brandHint(), question, repair);
}
private String brandHint() {
String b = identity == null ? null : identity.brandsId();
return (b == null || b.isBlank()) ? "本企业" : b;
}
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();
}
// AI 可用业务表/视图白名单(静态元数据,跨请求缓存);含 sBrandsId 列的表缓存。
private static volatile Set<String> ALLOWED_TABLES;
private static final ConcurrentHashMap<String, Boolean> BRAND_COL = new ConcurrentHashMap<>();
/**
* 单条 SELECT + 挡危险构造 + **表白名单**(架构 §9 的“视图白名单”)。
* 白名单 = AI 可用的业务表/视图(viw_* + 表单数据源 + 字段字典表),把凭证/权限/暂存等敏感表挡在外面
* (如 gdslogininfo、sysjurisdiction、ai_op_queue),防 NL2SQL 击穿权限。返回 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;
}
}
Statement stmt;
try {
stmt = CCJSqlParserUtil.parse(sql);
} catch (Exception e) {
return "SQL 解析失败";
}
if (!(stmt instanceof Select)) {
return "只允许 SELECT";
}
Set<String> allowed = allowedTables();
if (!allowed.isEmpty()) {
List<String> tables;
try {
tables = new TablesNamesFinder().getTableList(stmt);
} catch (Exception e) {
return "无法解析查询涉及的表";
}
for (String t : tables) {
if (!allowed.contains(normTable(t))) {
return "涉及不允许访问的表(" + normTable(t) + ")";
}
}
}
return null;
}
/** AI 可用表/视图集合(小写):所有 viw_* 视图 + 表单数据源 + 字段字典里出现的基础表。静态缓存。 */
private Set<String> allowedTables() {
Set<String> c = ALLOWED_TABLES;
if (c != null) {
return c;
}
Set<String> s = new HashSet<>();
try {
for (Map<String, Object> r : jdbc.queryForList(
"SELECT LOWER(TABLE_NAME) t FROM information_schema.VIEWS WHERE TABLE_SCHEMA=DATABASE() AND TABLE_NAME LIKE 'viw\\_%'")) {
s.add(String.valueOf(r.get("t")));
}
for (Map<String, Object> r : jdbc.queryForList(
"SELECT DISTINCT LOWER(sDataSource) t FROM viw_ai_useful_forms WHERE IFNULL(sDataSource,'')<>''")) {
s.add(String.valueOf(r.get("t")));
}
for (Map<String, Object> r : jdbc.queryForList(
"SELECT DISTINCT LOWER(sTable) t FROM viw_kg_field_dict WHERE IFNULL(sTable,'')<>''")) {
s.add(String.valueOf(r.get("t")));
}
} catch (Exception ignore) {
// 元数据不可用时返回空集 → 不启用白名单(保持可用),但仍有 SELECT-only + 禁止构造兜底。
}
ALLOWED_TABLES = s;
return s;
}
/** 规整表名:去反引号、去 schema 前缀、转小写。 */
private static String normTable(String t) {
if (t == null) {
return "";
}
String x = t.replace("`", "").trim();
int dot = x.lastIndexOf('.');
if (dot >= 0) {
x = x.substring(dot + 1);
}
return x.toLowerCase();
}
/**
* 租户注入(架构 §9):单表(无 JOIN)且该表含 sBrandsId 列、且身份带品牌时,追加
* {@code AND sBrandsId='<brand>'},把结果限定在本企业。视图(viw_*)通常已按品牌预筛,跳过。
* 任何异常都退回原 SQL(不因注入失败而阻断,白名单已是主要边界)。
*/
private String applyTenant(String sql) {
String brand = identity == null ? null : identity.brandsId();
if (brand == null || brand.isBlank()) {
return sql;
}
try {
Statement stmt = CCJSqlParserUtil.parse(sql);
if (!(stmt instanceof Select)) {
return sql;
}
Select select = (Select) stmt;
if (!(select.getSelectBody() instanceof PlainSelect)) {
return sql;
}
PlainSelect ps = (PlainSelect) select.getSelectBody();
if (ps.getJoins() != null && !ps.getJoins().isEmpty()) {
return sql; // 多表:交给白名单,不做注入
}
if (!(ps.getFromItem() instanceof Table)) {
return sql;
}
String table = normTable(((Table) ps.getFromItem()).getName());
if (table.startsWith("viw")) {
return sql; // 视图预筛,不注入
}
if (!hasBrandCol(table)) {
return sql;
}
EqualsTo eq = new EqualsTo();
eq.setLeftExpression(new Column("sBrandsId"));
eq.setRightExpression(new StringValue(brand));
ps.setWhere(ps.getWhere() == null ? eq : new AndExpression(ps.getWhere(), eq));
return select.toString();
} catch (Exception e) {
return sql;
}
}
/** 该基础表是否有 sBrandsId 列(缓存)。 */
private boolean hasBrandCol(String table) {
return BRAND_COL.computeIfAbsent(table, t -> {
try {
Integer n = jdbc.queryForObject(
"SELECT COUNT(*) FROM information_schema.COLUMNS WHERE TABLE_SCHEMA=DATABASE() " +
"AND TABLE_NAME=? AND COLUMN_NAME='sBrandsId'", Integer.class, t);
return n != null && n > 0;
} catch (Exception e) {
return false;
}
});
}
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();
}
}