using System.Text;
using System.Text.RegularExpressions;
namespace Admin.NET.Plugin.AiDOP.SmartOps;
///
/// KPI 配置 SQL 安全校验(多层、非仅正则)。真正的写保护由只读事务在 DB 层强制(KpiSqlReadOnlyExecutor);
/// 本校验器是意图闸门 + 纵深防御:单语句、起始 token 白名单、token 级 denylist、FROM 表白名单。
/// 顺序:去注释 → 单语句门 → 起始 token 白名单(SELECT/WITH) → denylist(token 边界) → FROM 表白名单。
///
public static class KpiSqlSecurityValidator
{
/// 允许查询的表前缀(中台内部表;租户谓词软限降级——限制 FROM 目标而非静态分析谓词)。
public static readonly string[] AllowedTablePrefixes =
{
"mdp_std_", "mdp_stg_", "dwd_", "ado_s9_kpi_value_", "ado_smart_ops_kpi_atomic"
};
// token 边界匹配的危险关键字(去注释后,按单词边界;避免子串误伤列名如 create_time)。
private static readonly string[] DenyKeywords =
{
"INSERT", "UPDATE", "DELETE", "MERGE", "REPLACE", "DROP", "ALTER", "CREATE", "TRUNCATE",
"GRANT", "REVOKE", "RENAME", "LOAD", "CALL", "EXEC", "EXECUTE", "SET", "LOCK", "UNLOCK",
"HANDLER", "PREPARE", "DEALLOCATE", "INTO", "OUTFILE", "DUMPFILE", "USE", "SHOW", "DESCRIBE",
"ATTACH", "PRAGMA", "COPY"
};
private static readonly HashSet StartTokenWhitelist = new(StringComparer.OrdinalIgnoreCase) { "SELECT", "WITH" };
public sealed class Result
{
public bool Ok { get; set; }
public string? ErrorCode { get; set; }
public string? ErrorMessage { get; set; }
public List ReferencedTables { get; } = new();
public static Result Fail(string code, string msg) => new() { Ok = false, ErrorCode = code, ErrorMessage = msg };
}
public static Result Validate(string? rawSql)
{
if (string.IsNullOrWhiteSpace(rawSql))
return Result.Fail("EMPTY_SQL", "SQL 不能为空");
// 1) 去注释:-- 行注释、# 行注释、/* */ 块注释(防止注释绕过关键字/多语句检测)
var sql = StripComments(rawSql).Trim();
if (sql.Length == 0)
return Result.Fail("EMPTY_SQL", "SQL 去注释后为空");
// 2) 单语句门:去尾部分号后,若仍含分号 → 多语句,拒
var trimmed = sql.TrimEnd().TrimEnd(';').TrimEnd();
if (trimmed.Contains(';'))
return Result.Fail("MULTI_STATEMENT", "只允许单条语句,禁止分号多语句");
if (trimmed.Length == 0)
return Result.Fail("EMPTY_SQL", "SQL 为空语句");
// token 化(按非标识符字符切分,保留 . 以便识别 schema.table)
var tokens = Tokenize(trimmed);
if (tokens.Count == 0)
return Result.Fail("EMPTY_SQL", "SQL 无有效 token");
// 3) 起始 token 白名单:必须以 SELECT 或 WITH 开头
if (!StartTokenWhitelist.Contains(tokens[0]))
return Result.Fail("NOT_SELECT", "只允许 SELECT 或 WITH ... SELECT 查询");
// WITH 需其后出现 SELECT 主语句
if (string.Equals(tokens[0], "WITH", StringComparison.OrdinalIgnoreCase)
&& !tokens.Any(t => string.Equals(t, "SELECT", StringComparison.OrdinalIgnoreCase)))
return Result.Fail("NOT_SELECT", "WITH 之后必须为 SELECT 查询");
// 4) token 级 denylist(单词边界匹配,非裸子串)
var upperTokens = new HashSet(tokens.Select(t => t.ToUpperInvariant()));
foreach (var kw in DenyKeywords)
{
if (upperTokens.Contains(kw))
return Result.Fail("FORBIDDEN_KEYWORD", $"禁止的关键字/操作:{kw}");
}
// 5) FROM 表白名单:提取 FROM/JOIN 后的表名,须命中允许前缀
var result = new Result { Ok = true };
var tables = ExtractTables(tokens);
result.ReferencedTables.AddRange(tables);
foreach (var tbl in tables)
{
var bare = tbl.Contains('.') ? tbl.Substring(tbl.LastIndexOf('.') + 1) : tbl;
bare = bare.Trim('`', '"', '[', ']');
if (!AllowedTablePrefixes.Any(p => bare.StartsWith(p, StringComparison.OrdinalIgnoreCase)))
return Result.Fail("TABLE_NOT_ALLOWED",
$"表 {tbl} 不在允许的中台数据集范围(仅 {string.Join("/", AllowedTablePrefixes)})");
}
return result;
}
private static string StripComments(string sql)
{
// 块注释 /* ... */(含跨行)
sql = Regex.Replace(sql, @"/\*.*?\*/", " ", RegexOptions.Singleline);
var sb = new StringBuilder();
foreach (var line in sql.Split('\n'))
{
var l = line;
var dash = l.IndexOf("--", StringComparison.Ordinal);
var hash = l.IndexOf('#');
var cut = -1;
if (dash >= 0) cut = dash;
if (hash >= 0 && (cut < 0 || hash < cut)) cut = hash;
if (cut >= 0) l = l.Substring(0, cut);
sb.Append(l).Append('\n');
}
return sb.ToString();
}
private static List Tokenize(string sql)
{
// 保留标识符字符 + . + @(参数);其余作分隔。字符串字面量整体折叠为占位,避免其中关键字误判。
var noStrings = Regex.Replace(sql, @"'([^'\\]|\\.)*'", " '' ");
noStrings = Regex.Replace(noStrings, "\"([^\"\\\\]|\\\\.)*\"", " \"\" ");
var matches = Regex.Matches(noStrings, @"[A-Za-z0-9_@\.\$]+");
return matches.Select(m => m.Value).Where(v => v.Length > 0).ToList();
}
private static List ExtractTables(List tokens)
{
var tables = new List();
for (var i = 0; i < tokens.Count - 1; i++)
{
var t = tokens[i].ToUpperInvariant();
if (t == "FROM" || t == "JOIN")
{
var next = tokens[i + 1];
// 跳过子查询开括号场景:Tokenize 已去括号,子查询里 FROM 会各自被扫描到;
// 若 next 是保留字(SELECT 等)说明是子查询,跳过表名收集。
if (StartTokenWhitelist.Contains(next) || string.Equals(next, "SELECT", StringComparison.OrdinalIgnoreCase))
continue;
tables.Add(next);
}
}
return tables;
}
}