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 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题 * (跨表汇总、按条件计数排名等)。 * *

安全栈:用 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> 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> 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; } 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", " "); if (s.length() > 30) { s = s.substring(0, 30) + "…"; } line.append(s).append(" | "); } sb.append(line).append("\n"); } return sb.toString(); } }