package com.trafficaudit.llmintegration.service; import cn.hutool.json.JSONObject; import lombok.extern.slf4j.Slf4j; import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.stereotype.Component; import javax.annotation.Resource; import java.util.ArrayList; import java.util.Arrays; import java.util.Collections; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Set; import java.util.TreeSet; import java.util.regex.Matcher; import java.util.regex.Pattern; /** * AI 对话查询数据库的安全执行器:只允许查询白名单表与白名单列, * 过滤条件由列名+值组成,避免任意 SQL 注入。 */ @Slf4j @Component public class DbQueryExecutor { @Resource private JdbcTemplate jdbcTemplate; public static final Set ALLOWED_TABLES = new TreeSet<>(); private static final Map TABLES = new LinkedHashMap<>(); private static final Pattern NUM_CMP = Pattern.compile("^(>=|<=|!=|>|<|=)(-?\\d+(\\.\\d+)?)$"); static { register("h2032_enterprise_monthly", "企业月报H203-2表(道路货物运输月度生产情况)", new String[]{"id", "report_period", "region_code", "enterprise_code", "enterprise_name", "unified_credit_code", "report_unit", "vehicle_total", "tons_total", "vehicle_tractor", "vehicle_trailer", "tons_trailer", "vehicle_container", "tons_container", "vehicle_whole", "tons_whole", "freight_total", "turnover_total", "freight_container", "freight_coal", "freight_oil_gas", "freight_crude_oil", "freight_metal_ore", "freight_iron_ore", "freight_building", "freight_grain", "avg_tonnage", "avg_distance", "unit_leader", "stats_leader", "contact_person", "contact_phone", "report_date", "verify_explanation", "report_notes"}, new String[]{"主键", "报表期", "所属地区代码", "企业代码", "企业名称", "统一社会信用代码", "填报单位", "车辆数合计", "标记吨位数合计", "牵引车数", "挂车数", "挂车吨位", "集装箱车数", "集装箱吨位", "整车数", "整车吨位", "货运量合计(吨)", "周转量合计(吨公里)", "货运量-集装箱", "货运量-煤炭及制品", "货运量-石油天然气", "货运量-原油", "货运量-金属矿石", "货运量-铁矿石", "货运量-建材", "货运量-粮食", "平均吨位", "平均运距(公里)", "单位负责人", "统计负责人", "填表人", "联系电话", "报出日期", "企业核实解释", "备注"}, new String[]{"id", "report_period", "enterprise_name", "vehicle_total", "tons_total", "freight_total", "turnover_total", "avg_distance", "verify_explanation"}); register("transport_auth_vehicle", "运政车辆", new String[]{"id", "report_period", "enterprise_name", "tractor_count", "trailer_count", "other_count", "trailer_tons", "other_tons"}, new String[]{"主键", "报表期", "企业名称", "牵引车数", "挂车数", "其他车数", "挂车吨位", "其他吨位"}, new String[]{"id", "report_period", "enterprise_name", "tractor_count", "trailer_count", "other_count", "trailer_tons", "other_tons"}); register("vehicle_track_mileage", "车辆轨迹里程", new String[]{"id", "report_period", "enterprise_name", "monthly_mileage", "tracked_vehicles"}, new String[]{"主键", "报表期", "企业名称", "月度行驶里程(公里)", "有轨迹车辆数"}, new String[]{"id", "report_period", "enterprise_name", "monthly_mileage", "tracked_vehicles"}); register("scale_split_transport", "规上规下拆分运输量", new String[]{"id", "report_period", "period_type", "region_name", "above_scale_freight", "above_scale_turnover", "above_scale_rank", "above_scale_yoy", "below_scale_freight", "below_scale_turnover", "below_scale_rank", "below_scale_yoy", "total_freight", "total_turnover", "total_rank", "total_yoy"}, new String[]{"主键", "报表期", "口径(MONTH当月/CUMULATIVE累计)", "市州名称(含全省)", "规上货运量", "规上周转量", "规上排名", "规上同比", "规下货运量", "规下周转量", "规下排名", "规下同比", "合计货运量", "合计周转量", "合计排名", "合计同比"}, new String[]{"report_period", "period_type", "region_name", "above_scale_freight", "above_scale_turnover", "below_scale_freight", "below_scale_turnover", "total_freight", "total_turnover"}); List fcols = new ArrayList<>(); List flabels = new ArrayList<>(); fcols.add("id"); flabels.add("主键"); fcols.add("report_period"); flabels.add("报表期"); fcols.add("region_name"); flabels.add("市州名称"); for (int m = 1; m <= 12; m++) { fcols.add(String.format("freight_m%02d", m)); flabels.add("今年" + m + "月货运量"); } for (int m = 1; m <= 12; m++) { fcols.add(String.format("last_freight_m%02d", m)); flabels.add("去年" + m + "月货运量"); } for (int m = 1; m <= 12; m++) { fcols.add(String.format("turnover_m%02d", m)); flabels.add("今年" + m + "月周转量"); } for (int m = 1; m <= 12; m++) { fcols.add(String.format("last_turnover_m%02d", m)); flabels.add("去年" + m + "月周转量"); } register("freight_turnover_import", "货运量周转量(月度导入)", fcols.toArray(new String[0]), flabels.toArray(new String[0]), new String[]{"report_period", "region_name", "freight_m01", "freight_m02", "freight_m03", "freight_m04", "freight_m05", "freight_m06", "freight_m07", "freight_m08", "freight_m09", "freight_m10", "freight_m11", "freight_m12", "turnover_m01", "turnover_m12"}); register("h2031_enterprise_monthly", "公路旅客月报H203-1表(道路旅客运输月度生产情况)", new String[]{"id", "report_period", "region_code", "enterprise_code", "enterprise_name", "vehicle_total", "seat_total", "passenger_total", "turnover_total", "avg_distance_total", "verify_explanation"}, new String[]{"主键", "报表期", "所属地区代码", "企业代码", "企业名称", "车辆数", "载客位数", "客运量(万人)", "旅客周转量(万人公里)", "平均运距(公里)", "企业核实解释"}, new String[]{"id", "report_period", "enterprise_name", "vehicle_total", "seat_total", "passenger_total", "turnover_total", "avg_distance_total", "verify_explanation"}); register("city_bus_monthly", "城市公交月度运营情况(企业级,含轨道/轮渡字段)", new String[]{"id", "report_period", "region_code", "city", "enterprise_name", "op_vehicles", "passenger_volume", "turnover", "avg_distance", "passenger_chengxiang", "turnover_chengxiang", "verify_explanation"}, new String[]{"主键", "报表期", "所属地区代码", "市州", "企业名称", "运营车数", "客运量(万人次)", "旅客周转量(万人公里)", "平均运距(公里)", "城际城乡客运量", "城际城乡周转量", "企业核实解释"}, new String[]{"id", "report_period", "enterprise_name", "op_vehicles", "passenger_volume", "turnover", "avg_distance", "verify_explanation"}); register("city_taxi_monthly", "巡游出租汽车运营服务情况月报(市州级)", new String[]{"id", "report_period", "region_code", "city", "trip_total", "passenger_volume", "turnover", "op_vehicles", "avg_distance", "verify_explanation"}, new String[]{"主键", "报表期", "所属地区代码", "市州", "载客车次总数", "客运量(万人次)", "旅客周转量(万人公里)", "运营车辆数", "平均运距(公里)", "企业核实解释"}, new String[]{"id", "report_period", "city", "trip_total", "passenger_volume", "turnover", "op_vehicles", "avg_distance", "verify_explanation"}); register("audit_result", "审核结果", new String[]{"id", "rule_id", "report_id", "enterprise_code", "report_period", "actual_value", "threshold_value", "deviation", "status", "review_comment", "created_at"}, new String[]{"主键", "规则ID", "上报记录ID", "企业代码", "报表期", "实际值", "阈值/对比值", "偏差说明", "状态(PENDING待处理/CONFIRMED确认/IGNORED忽略)", "审核意见", "创建时间"}, new String[]{"id", "report_id", "enterprise_code", "report_period", "actual_value", "threshold_value", "deviation", "status"}); register("audit_rule", "审核规则", new String[]{"id", "rule_name", "rule_code", "description", "check_field", "compare_type", "threshold", "alert_level", "report_type", "is_enabled"}, new String[]{"主键", "规则名称", "规则编码", "规则描述", "检查字段", "比较方式", "阈值", "预警级别", "报表类型", "是否启用"}, new String[]{"id", "rule_name", "rule_code", "description", "check_field", "compare_type", "threshold", "alert_level"}); } private static void register(String table, String desc, String[] cols, String[] labels, String[] defaults) { Map labelMap = new LinkedHashMap<>(); for (int i = 0; i < cols.length && i < labels.length; i++) { labelMap.put(cols[i], labels[i]); } TABLES.put(table, new TableMeta(table, desc, cols, labelMap, defaults)); ALLOWED_TABLES.add(table); } /** 供系统提示词使用的完整表结构说明 */ public String schemaText() { StringBuilder sb = new StringBuilder(); for (TableMeta meta : TABLES.values()) { sb.append(meta.table).append("【").append(meta.desc).append("】: "); List parts = new ArrayList<>(); for (String c : meta.cols) { parts.add(snakeToCamel(c) + "(" + meta.labels.getOrDefault(c, snakeToCamel(c)) + ")"); } sb.append(String.join(", ", parts)).append("\n"); } return sb.toString(); } public String execute(String table, Object columnsParam, Object filtersParam, String orderBy, Integer limit, String defaultPeriod) { JSONObject out = new JSONObject(); try { TableMeta meta = TABLES.get(table); if (meta == null) { out.set("success", false); out.set("error", "表名不合法,可用表: " + ALLOWED_TABLES); return out.toString(); } List cols = resolveCols(meta, columnsParam); Map filters = filtersParam instanceof Map ? (Map) filtersParam : new LinkedHashMap<>(); WhereClause w = buildWhere(meta, filters, defaultPeriod); StringBuilder sql = new StringBuilder("SELECT ").append(String.join(", ", cols)) .append(" FROM ").append(meta.table); if (!w.sql.isEmpty()) { sql.append(" WHERE ").append(w.sql); } sql.append(" ORDER BY ").append(parseOrder(meta, orderBy)); int lim = limit == null ? 20 : Math.min(Math.max(limit, 1), 100); sql.append(" LIMIT ").append(lim); List> rows = jdbcTemplate.queryForList(sql.toString(), w.args.toArray()); out.set("success", true); out.set("rows", rows); out.set("count", rows.size()); out.set("total", countTotal(meta, filters, defaultPeriod)); return out.toString(); } catch (Exception ex) { log.warn("query_database tool error: {}", ex.getMessage()); out.set("success", false); out.set("error", ex.getMessage() == null ? "查询失败" : ex.getMessage()); return out.toString(); } } private List resolveCols(TableMeta meta, Object columnsParam) { if (columnsParam == null) { return aliased(meta.defaults); } if (columnsParam instanceof String && "*".equals(columnsParam)) { return aliased(meta.cols); } if (columnsParam instanceof List) { List out = new ArrayList<>(); for (Object o : (List) columnsParam) { String c = normalizeCol(String.valueOf(o)); if (!meta.cols.contains(c)) { throw new IllegalArgumentException("列名不合法: " + o + ",该表可用列: " + camelList(meta.cols)); } out.add(c); } return aliased(out.isEmpty() ? meta.defaults : out); } return aliased(meta.defaults); } private List aliased(List cols) { List out = new ArrayList<>(); for (String c : cols) { out.add(c + " AS " + snakeToCamel(c)); } return out; } private String camelList(List cols) { List parts = new ArrayList<>(); for (String c : cols) { parts.add(snakeToCamel(c)); } return String.join(", ", parts); } private WhereClause buildWhere(TableMeta meta, Map filters, String defaultPeriod) { List conds = new ArrayList<>(); List args = new ArrayList<>(); boolean hasPeriod = meta.cols.contains("report_period"); if (defaultPeriod != null && !defaultPeriod.isEmpty() && hasPeriod && !filters.containsKey("report_period")) { conds.add("report_period = ?"); args.add(defaultPeriod); } for (Map.Entry e : filters.entrySet()) { String col = normalizeCol(String.valueOf(e.getKey())); if (!meta.cols.contains(col)) { throw new IllegalArgumentException("列名不合法: " + e.getKey() + ",该表可用列: " + camelList(meta.cols)); } Object v = e.getValue(); if (v instanceof List) { List list = (List) v; if (list.isEmpty()) { continue; } conds.add(col + " IN (" + String.join(",", Collections.nCopies(list.size(), "?")) + ")"); args.addAll(list); } else if (v instanceof Number || v instanceof Boolean) { conds.add(col + " = ?"); args.add(v); } else { String s = String.valueOf(v).trim(); Matcher m = NUM_CMP.matcher(s); if (m.matches()) { conds.add(col + " " + m.group(1) + " ?"); args.add(Double.parseDouble(m.group(2))); } else if (!s.isEmpty()) { conds.add(col + " LIKE ?"); args.add("%" + s + "%"); } } } return new WhereClause(String.join(" AND ", conds), args); } private String parseOrder(TableMeta meta, String orderBy) { if (orderBy == null || orderBy.trim().isEmpty()) { return "id DESC"; } String s = orderBy.trim(); boolean desc = false; if (s.toUpperCase().endsWith(" DESC")) { desc = true; s = s.substring(0, s.length() - 5).trim(); } else if (s.toUpperCase().endsWith(" ASC")) { s = s.substring(0, s.length() - 4).trim(); } String col = normalizeCol(s); if (!meta.cols.contains(col)) { throw new IllegalArgumentException("排序列不合法: " + orderBy); } return col + (desc ? " DESC" : " ASC"); } private long countTotal(TableMeta meta, Map filters, String defaultPeriod) { WhereClause w = buildWhere(meta, filters, defaultPeriod); String sql = "SELECT COUNT(*) FROM " + meta.table; if (!w.sql.isEmpty()) { sql += " WHERE " + w.sql; } Long c = jdbcTemplate.queryForObject(sql, Long.class, w.args.toArray()); return c == null ? 0 : c; } private static String normalizeCol(String name) { StringBuilder sb = new StringBuilder(); for (char c : name.toCharArray()) { if (Character.isUpperCase(c)) { sb.append('_').append(Character.toLowerCase(c)); } else { sb.append(c); } } return sb.toString(); } private static String snakeToCamel(String s) { StringBuilder sb = new StringBuilder(); boolean up = false; for (char c : s.toCharArray()) { if (c == '_') { up = true; continue; } sb.append(up ? Character.toUpperCase(c) : c); up = false; } return sb.toString(); } private static class TableMeta { final String table; final String desc; final List cols; final Map labels; final List defaults; TableMeta(String table, String desc, String[] cols, Map labels, String[] defaults) { this.table = table; this.desc = desc; this.cols = Arrays.asList(cols); this.labels = labels; this.defaults = Arrays.asList(defaults); } } private static class WhereClause { final String sql; final List args; WhereClause(String sql, List args) { this.sql = sql; this.args = args; } } }