QueryTool.java 18.5 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 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416
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.chat.ChatModel;
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 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题
 * (跨表汇总、按条件计数排名等)。
 *
 * <p>安全栈(架构 §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 全部入审计。
 *
 * <p>尚未落地的加固项:多表 JOIN 只靠 prompt 提示带租户过滤(AST 注入只覆盖单表)、
 * 以及独立的只读 MySQL 账号(当前与应用共用连接池)。
 */
public class QueryTool {

    private final ChatModel sqlModel;
    private final JdbcTemplate jdbc;
    private final AuditService audit;
    private final AgentIdentity identity;
    private final FormResolverService resolver;

    public QueryTool(ChatModel 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<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 用英文列名或序号,**绝不要用中文做别名**;
                **给用户看的是名称不是内部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<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;
    }

    // 裸主键ID(业务 sId 通常为 15+ 位数字串)。仅对严格匹配的长数字尝试解析,普通数字(数量/金额)不受影响。
    private static final Pattern LONG_ID = Pattern.compile("\\d{15,}");
    private static volatile Map<String, String> NAME_TABLES; // 主表 -> 名称字段(静态缓存)
    private final Map<String, String> idNameCache = new ConcurrentHashMap<>(); // 本次请求内的 id->名称 缓存

    /** 被外键引用的主表集合及其名称字段(客户/产品/物料/供应商…)。用于把裸ID回解析成名称。 */
    private Map<String, String> nameTables() {
        Map<String, String> c = NAME_TABLES;
        if (c != null) {
            return c;
        }
        Map<String, String> m = new LinkedHashMap<>();
        try {
            for (Map<String, Object> 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<String, String> e : nameTables().entrySet()) {
            try {
                List<Map<String, Object>> 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<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", " ");
                // 把裸的长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();
    }
}