QueryTool.java 14.9 KB
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();
    }
}