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