pg sql解析器 给所有语句增加多组织查询sql

nuget

    <PackageReference Include="Npgquery" Version="1.1.0" />

解析器

/// <summary>
/// 基于 Npgquery 解析 PostgreSQL SQL,自动为业务表添加 orgid 过滤条件
/// </summary>
public class PostgreSqlOrgFilter
{
    // 白名单:没有 orgid 字段的表(全局表、字典表、配置表等)
    private static readonly HashSet<string> SkipTables = new(StringComparer.OrdinalIgnoreCase)
    {
        "sys_dict",
        "app_config",
        "user_org_roles",
        "migrations",
        "__efmigrationshistory"
    };

    /// <summary>
    /// 为 SQL 中所有业务表添加 orgid = @OrgId 条件
    /// </summary>
    public string AddOrgFilter(string sql, string orgIdParam = "@OrgId")
    {
        if (string.IsNullOrWhiteSpace(sql))
            return sql;

        // 1. 解析 SQL,提取表名
        var tables = ExtractTableNames(sql);
        if (tables.Count == 0)
            return sql;

        // 2. 过滤掉白名单表
        var orgTables = tables
            .Where(t => !SkipTables.Contains(t.Table))
            .Distinct()
            .ToList();

        if (orgTables.Count == 0)
            return sql;

        // 3. 构建要追加的条件
        var conditions = orgTables
            .Select(t => $"{t.Alias ?? t.Table}.orgid = {orgIdParam}")
            .ToList();

        var conditionStr = string.Join(" AND ", conditions);

        // 4. 找到插入位置(WHERE / GROUP BY / ORDER BY / LIMIT / 结尾)
        int insertPos = FindInsertPosition(sql);

        // 5. 判断是否需要加 WHERE 还是 AND
        bool hasWhere = Regex.IsMatch(sql, @"\bWHERE\b", RegexOptions.IgnoreCase);

        string appendSql;
        if (hasWhere)
        {
            appendSql = " AND " + conditionStr;
        }
        else
        {
            appendSql = " WHERE " + conditionStr;
        }

        return sql.Insert(insertPos, appendSql);
    }

    /// <summary>
    /// 使用 Npgquery 解析 SQL,提取所有表名和别名
    /// </summary>
    private List<(string Table, string Alias)> ExtractTableNames(string sql)
    {
        var result = Parser.QuickParse(sql);
        if (!result.IsSuccess)
        {
            // 解析失败时,退回到正则提取(兜底)
            return FallbackExtractTableNames(sql);
        }

        var tables = new List<(string, string)>();
        var root = result.ParseTree!.RootElement;
        WalkRangeVars(root, tables);
        return tables;
    }

    /// <summary>
    /// 递归遍历 AST,提取所有 RangeVar(表引用)
    /// </summary>
    private void WalkRangeVars(JsonElement node, List<(string, string)> tables)
    {
        if (node.ValueKind == JsonValueKind.Object)
        {
            // 检查是否是 RangeVar
            if (node.TryGetProperty("RangeVar", out var rangeVar))
            {
                string table = rangeVar.TryGetProperty("relname", out var rel) ? rel.GetString() : null;
                string alias = null;

                if (rangeVar.TryGetProperty("alias", out var aliasObj) &&
                    aliasObj.TryGetProperty("aliasname", out var aliasName))
                {
                    alias = aliasName.GetString();
                }

                if (table != null)
                {
                    tables.Add((table, alias ?? table));
                }
            }

            // 递归遍历所有属性
            foreach (var prop in node.EnumerateObject())
            {
                WalkRangeVars(prop.Value, tables);
            }
        }
        else if (node.ValueKind == JsonValueKind.Array)
        {
            foreach (var item in node.EnumerateArray())
            {
                WalkRangeVars(item, tables);
            }
        }
    }

    /// <summary>
    /// 正则兜底:提取 FROM/JOIN 后的表名
    /// </summary>
    private List<(string Table, string Alias)> FallbackExtractTableNames(string sql)
    {
        var tables = new List<(string, string)>();
        var pattern = new Regex(
            @"\b(?:FROM|JOIN)\s+""?(\w+)""?(?:\s+(?:AS\s+)?""?(\w+)""?)?",
            RegexOptions.IgnoreCase | RegexOptions.Multiline);

        foreach (Match match in pattern.Matches(sql))
        {
            var table = match.Groups[1].Value;
            var alias = match.Groups[2].Success ? match.Groups[2].Value : table;
            tables.Add((table, alias));
        }

        return tables;
    }

    /// <summary>
    /// 找到合适的插入位置(WHERE / GROUP BY / ORDER BY / LIMIT / 结尾)
    /// </summary>
    private int FindInsertPosition(string sql)
    {
        int pos = sql.Length;

        // 按优先级找最后一个关键字的位置
        var keywords = new[] { "GROUP BY", "ORDER BY", "LIMIT", "OFFSET", "FOR UPDATE", "FOR SHARE" };
        foreach (var kw in keywords)
        {
            int idx = sql.IndexOf(kw, StringComparison.OrdinalIgnoreCase);
            if (idx > -1 && idx < pos)
            {
                pos = idx;
            }
        }

        // 如果有 UNION / INTERSECT / EXCEPT,只在第一个 SELECT 块内加
        int unionIdx = sql.IndexOf("UNION", StringComparison.OrdinalIgnoreCase);
        if (unionIdx > -1 && unionIdx < pos)
        {
            pos = unionIdx;
        }

        return pos;
    }
}
posted @ 2026-09-28 14:25  Hey,Coder!  阅读(2)  评论(0)    收藏  举报