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<String> ALLOWED_TABLES = new TreeSet<>();
|
|
private static final Map<String, TableMeta> 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<String> fcols = new ArrayList<>();
|
List<String> 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("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<String, String> 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<String> 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<String> cols = resolveCols(meta, columnsParam);
|
Map<String, Object> filters = filtersParam instanceof Map
|
? (Map<String, Object>) 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<Map<String, Object>> 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<String> 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<String> 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<String> aliased(List<String> cols) {
|
List<String> out = new ArrayList<>();
|
for (String c : cols) {
|
out.add(c + " AS " + snakeToCamel(c));
|
}
|
return out;
|
}
|
|
private String camelList(List<String> cols) {
|
List<String> parts = new ArrayList<>();
|
for (String c : cols) {
|
parts.add(snakeToCamel(c));
|
}
|
return String.join(", ", parts);
|
}
|
|
private WhereClause buildWhere(TableMeta meta, Map<String, Object> filters, String defaultPeriod) {
|
List<String> conds = new ArrayList<>();
|
List<Object> 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<String, Object> 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<String, Object> 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<String> cols;
|
final Map<String, String> labels;
|
final List<String> defaults;
|
|
TableMeta(String table, String desc, String[] cols, Map<String, String> 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<Object> args;
|
|
WhereClause(String sql, List<Object> args) {
|
this.sql = sql;
|
this.args = args;
|
}
|
}
|
}
|