QueryTool.java 9.04 KB
package com.xly.tool;

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.parser.CCJSqlParserUtil;
import net.sf.jsqlparser.statement.Statement;
import net.sf.jsqlparser.statement.select.Select;
import org.springframework.beans.factory.annotation.Qualifier;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.stereotype.Component;

import java.util.List;
import java.util.Map;

/**
 * Query 工具:**只读 SQL 兜底** —— 回答没有现成表单/记录能直接答的临时统计/分析问题
 * (跨表汇总、按条件计数排名等)。
 *
 * <p>安全栈:用 coder 模型据 KG 字段字典接地生成 SQL → jsqlparser 强制**单条 SELECT** →
 * 挡 {@code INTO OUTFILE / LOAD_FILE / information_schema / SLEEP / BENCHMARK} 与多语句 →
 * 强制 LIMIT。(本地单品牌,租户注入留作生产加固;见架构 §9。)SQL 入审计。
 */
@Component
public class QueryTool {

    private final OllamaChatModel sqlModel;
    private final JdbcTemplate jdbc;
    private final AuditService audit;

    public QueryTool(@Qualifier("sqlChatModel") OllamaChatModel sqlModel, JdbcTemplate jdbc, AuditService audit) {
        this.sqlModel = sqlModel;
        this.jdbc = jdbc;
        this.audit = audit;
    }

    @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(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);
                **列别名一律用英文**(如 cnt、total、name),ORDER BY 用英文列名或序号,**绝不要用中文做别名**;
                表名、列名一律用上面给定的英文名。**只输出 SQL 本身**,不要解释、不要 markdown 代码围栏。
                问题:%s%s
                """.formatted(hint, question, repair);
    }

    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();
    }

    /** 单条 SELECT + 挡危险构造。返回 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;
            }
        }
        try {
            Statement stmt = CCJSqlParserUtil.parse(sql);
            if (!(stmt instanceof Select)) {
                return "只允许 SELECT";
            }
        } catch (Exception e) {
            return "SQL 解析失败";
        }
        return null;
    }

    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();
    }
}