package com.xly.tool; import com.xly.agent.AgentIdentity; import com.xly.service.AuditService; import com.xly.service.FormResolverService; 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.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; import java.util.regex.Pattern; /** * Query 工具:**只读 SQL 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题 * (跨表汇总、按条件计数排名等)。 * *

安全栈(架构 §9):用 coder 模型据 KG 字段字典接地生成 SQL → jsqlparser 强制**单条 SELECT** → * 挡 {@code INTO OUTFILE / LOAD_FILE / information_schema / SLEEP / BENCHMARK} 与多语句 → * **表白名单**({@link #allowedTables()}:viw_* + 表单数据源 + 字段字典表,挡掉 gdslogininfo / * sysjurisdiction / ai_op_queue 等敏感表) → **AST 注入租户谓词**({@link #applyTenant},单表查询按 * {@code sBrandsId}) → 强制 LIMIT。SQL 全部入审计。 * *

尚未落地的加固项:多表 JOIN 只靠 prompt 提示带租户过滤(AST 注入只覆盖单表)、 * 以及独立的只读 MySQL 账号(当前与应用共用连接池)。 */ public class QueryTool { private final OllamaChatModel sqlModel; private final JdbcTemplate jdbc; private final AuditService audit; private final AgentIdentity identity; private final FormResolverService resolver; public QueryTool(OllamaChatModel sqlModel, JdbcTemplate jdbc, AuditService audit, AgentIdentity identity, FormResolverService resolver) { this.sqlModel = sqlModel; this.jdbc = jdbc; this.audit = audit; this.identity = identity; this.resolver = resolver; } @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> 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 用英文列名或序号,**绝不要用中文做别名**; **给用户看的是名称不是内部ID**:若按某实体(产品/客户/物料/供应商)分组或排名,必须 JOIN 该实体主表、 在结果里返回它的**名称列**(如 eleproduct.sProductName、elecustomer.sCustomerName),不要只返回 *Id 列; 表名、列名一律用上面给定的英文名。**只输出 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> 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 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 ALLOWED_TABLES; private static final ConcurrentHashMap 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 allowed = allowedTables(); if (!allowed.isEmpty()) { List 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 allowedTables() { Set c = ALLOWED_TABLES; if (c != null) { return c; } Set s = new HashSet<>(); try { for (Map 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 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 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=''},把结果限定在本企业。视图(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; } // 裸主键ID(业务 sId 通常为 15+ 位数字串)。仅对严格匹配的长数字尝试解析,普通数字(数量/金额)不受影响。 private static final Pattern LONG_ID = Pattern.compile("\\d{15,}"); private static volatile Map NAME_TABLES; // 主表 -> 名称字段(静态缓存) private final Map idNameCache = new ConcurrentHashMap<>(); // 本次请求内的 id->名称 缓存 /** 被外键引用的主表集合及其名称字段(客户/产品/物料/供应商…)。用于把裸ID回解析成名称。 */ private Map nameTables() { Map c = NAME_TABLES; if (c != null) { return c; } Map m = new LinkedHashMap<>(); try { for (Map r : jdbc.queryForList( "SELECT DISTINCT sFkTable FROM viw_kg_field_dict WHERE IFNULL(sFkTable,'')<>'' " + "AND sFkTable NOT LIKE 'viw%'")) { String t = String.valueOf(r.get("sFkTable")); if (t == null || t.isBlank()) continue; String nf = resolver.resolveNameField(t); if (nf != null && !nf.isBlank()) { m.put(t, nf); } } } catch (Exception ignore) { } NAME_TABLES = m; return m; } /** 长ID → 名称(在被引用主表里按 sId 命中即返回;找不到返回 null,保留原ID)。 */ private String resolveId(String val) { String cached = idNameCache.get(val); if (cached != null) { return cached.isEmpty() ? null : cached; } for (Map.Entry e : nameTables().entrySet()) { try { List> r = jdbc.queryForList( "SELECT `" + e.getValue() + "` nm FROM `" + e.getKey() + "` WHERE sId=? LIMIT 1", val); if (!r.isEmpty()) { Object nm = r.get(0).get("nm"); if (nm != null && !String.valueOf(nm).isBlank()) { String s = String.valueOf(nm); idNameCache.put(val, s); return s; } } } catch (Exception ignore) { } } idNameCache.put(val, ""); return null; } private String formatRows(List> rows) { if (rows.isEmpty()) { return "查询完成,没有匹配的数据。"; } List 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 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", " "); // 把裸的长ID(如产品/客户 sId)就地解析成名称,给用户看名字不是ID if (LONG_ID.matcher(s).matches()) { String nm = resolveId(s); if (nm != null) s = nm; } if (s.length() > 30) { s = s.substring(0, 30) + "…"; } line.append(s).append(" | "); } sb.append(line).append("\n"); } return sb.toString(); } }