package com.l.tracecd.service; import com.l.tracecd.constant.Constants; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.stereotype.Service; /** * SQL 校验服务 * 校验 LLM 生成的 SQL 语句合法性,防止注入和危险操作 */ @Service public class SqlValidationService { private static final Logger log = LoggerFactory.getLogger(SqlValidationService.class); /** INSERT 语句禁止的关键字 */ private static final String[] INSERT_DANGEROUS = { "DROP", "DELETE", "UPDATE", "ALTER", "TRUNCATE", "CREATE", "EXEC", "EXECUTE", "UNION", "--", "/*", ";", "GRANT", "REVOKE" }; /** SELECT 语句禁止的关键字 */ private static final String[] SELECT_DANGEROUS = { "DROP", "DELETE", "UPDATE", "ALTER", "TRUNCATE", "CREATE", "INSERT", "EXEC", "EXECUTE", "--", "/*", "GRANT", "REVOKE" }; /** * 校验录入场景的 INSERT 语句 * * @param sql 待校验的 SQL * @return 清理后的 SQL * @throws SqlValidationException 校验失败 */ public String validateInsert(String sql) throws SqlValidationException { if (sql == null || sql.isBlank()) { throw new SqlValidationException("SQL 语句为空"); } String cleaned = cleanSql(sql); String upper = cleaned.toUpperCase().replaceAll("\\s+", " ").trim(); // 必须以 INSERT INTO 开头,且目标表为 t_daily_record if (!upper.startsWith("INSERT INTO") || !upper.contains(Constants.TABLE_DAILY_RECORD.toUpperCase())) { throw new SqlValidationException("SQL 必须是 INSERT INTO " + Constants.TABLE_DAILY_RECORD + " 语句"); } // 检查 VALUES 关键字 if (!upper.contains("VALUES")) { throw new SqlValidationException("INSERT 语句必须包含 VALUES"); } // 检查危险关键字 checkDangerous(upper, INSERT_DANGEROUS); log.debug("INSERT SQL 校验通过: {}", cleaned); return cleaned; } /** * 校验查询场景的 SELECT 语句 * * @param sql 待校验的 SQL * @return 清理后的 SQL * @throws SqlValidationException 校验失败 */ public String validateSelect(String sql) throws SqlValidationException { if (sql == null || sql.isBlank()) { throw new SqlValidationException("SQL 语句为空"); } String cleaned = cleanSql(sql); String upper = cleaned.toUpperCase().replaceAll("\\s+", " ").trim(); // 必须以 SELECT 开头 if (!upper.startsWith("SELECT")) { throw new SqlValidationException("查询只允许 SELECT 语句"); } // 必须包含 t_daily_record 表 if (!upper.contains(Constants.TABLE_DAILY_RECORD.toUpperCase())) { throw new SqlValidationException("只允许查询 " + Constants.TABLE_DAILY_RECORD + " 表"); } // 检查危险关键字 checkDangerous(upper, SELECT_DANGEROUS); log.debug("SELECT SQL 校验通过: {}", cleaned); return cleaned; } /** * 清理 SQL:去除首尾空白、末尾分号、markdown 代码块标记 */ private String cleanSql(String sql) { String cleaned = sql.trim(); // 去除 markdown 代码块 cleaned = cleaned.replaceAll("^```sql\\s*", "").replaceAll("^```\\s*", ""); cleaned = cleaned.replaceAll("```$", ""); // 去除末尾分号 cleaned = cleaned.replaceAll(";\\s*$", ""); return cleaned.trim(); } /** * 检查是否包含危险关键字 */ private void checkDangerous(String upperSql, String[] dangerous) throws SqlValidationException { for (String keyword : dangerous) { if (upperSql.contains(keyword)) { throw new SqlValidationException("SQL 包含非法关键字: " + keyword); } } } /** * SQL 校验异常 */ public static class SqlValidationException extends Exception { public SqlValidationException(String message) { super(message); } } }