TenantDatabaseBootstrapper.cs 7.4 KB
using System.Text.RegularExpressions;
using SqlSugar;
using Volo.Abp;
using Yi.Framework.SqlSugarCore.Abstractions;

namespace FoodLabeling.Th.Application.MultiTenancy;

/// <summary>
/// 租户业务库同步建库与 GRANT(建库候选优先级与 TenantService.CodeFirstForTenantAsync 一致)。
/// </summary>
public static class TenantDatabaseBootstrapper
{
    private static readonly Regex SafeDatabaseNameRegex =
        new(@"^[a-zA-Z0-9][a-zA-Z0-9_\-]{0,63}$", RegexOptions.Compiled);

    /// <summary>
    /// 同步执行 CREATE DATABASE IF NOT EXISTS 并对业务账号 GRANT;失败抛 <see cref="UserFriendlyException"/>。
    /// </summary>
    public static void EnsureDatabaseCreated(
        DbConnOptions dbConnOptions,
        DbType dbType,
        string tenantConnectionString,
        string databaseName)
    {
        if (string.IsNullOrWhiteSpace(databaseName) || !SafeDatabaseNameRegex.IsMatch(databaseName))
        {
            throw new UserFriendlyException("租户连接串中的数据库名无效,无法建库");
        }

        var createCandidates = BuildCreateDatabaseCandidates(dbConnOptions, tenantConnectionString);
        if (createCandidates.Count == 0)
        {
            throw new UserFriendlyException("未配置可用的数据库连接串,无法创建租户库");
        }

        var createErrors = new List<string>();
        string? privilegedConnectionString = null;
        foreach (var candidate in createCandidates)
        {
            try
            {
                ExecuteCreateDatabase(dbType, candidate.ConnectionString, databaseName);
                privilegedConnectionString = candidate.ConnectionString;
                break;
            }
            catch (Exception ex)
            {
                createErrors.Add($"{candidate.Label}: {GetRootMessage(ex)}");
            }
        }

        if (privilegedConnectionString == null)
        {
            var userId = ExtractConnectionValue(tenantConnectionString, "uid") ?? "netteam";
            throw new UserFriendlyException(
                $"创建租户库失败(库={databaseName})。尝试结果:{string.Join(" | ", createErrors)}。" +
                $"说明:MySQL 报 Access denied to database 新建库名时,通常是账号没有 CREATE 权限(与能否连上主库无关)。" +
                $"请在 appsettings 的 DbConnOptions.AdminConnectionString 配置高权限账号连接串后重试;或手动执行:" +
                $"CREATE DATABASE IF NOT EXISTS `{databaseName}` DEFAULT CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; " +
                $"GRANT ALL PRIVILEGES ON `{databaseName}`.* TO '{userId}'@'%'; " +
                $"再调用 initialize-tenant-database 补跑建表与 Seed");
        }

        TryGrantTenantDatabaseAccess(
            dbType,
            privilegedConnectionString,
            databaseName,
            tenantConnectionString);
    }

    /// <summary>
    /// 建库连接候选(按优先级,互不相同才尝试下一条)。
    /// </summary>
    internal static List<(string Label, string ConnectionString)> BuildCreateDatabaseCandidates(
        DbConnOptions dbConnOptions,
        string tenantConnectionString)
    {
        var createCandidates = new List<(string Label, string ConnectionString)>();
        void AddCandidate(string label, string? conn)
        {
            if (string.IsNullOrWhiteSpace(conn))
            {
                return;
            }

            if (createCandidates.Any(x =>
                    string.Equals(x.ConnectionString, conn, StringComparison.OrdinalIgnoreCase)))
            {
                return;
            }

            createCandidates.Add((label, conn));
        }

        AddCandidate("AdminConnectionString", dbConnOptions.AdminConnectionString);
        AddCandidate("DbConnOptions.Url(主库)", dbConnOptions.Url);
        AddCandidate(
            "实例级(无database)",
            BuildServerLevelConnectionString(
                !string.IsNullOrWhiteSpace(dbConnOptions.Url)
                    ? dbConnOptions.Url!
                    : tenantConnectionString));

        return createCandidates;
    }

    private static void ExecuteCreateDatabase(DbType dbType, string connectionString, string databaseName)
    {
        using var createDb = new SqlSugarClient(new ConnectionConfig
        {
            ConfigId = $"tenant-create-db-{databaseName}-{Guid.NewGuid():N}",
            DbType = dbType,
            ConnectionString = connectionString,
            IsAutoCloseConnection = true
        });

        createDb.Ado.ExecuteCommand(
            $"CREATE DATABASE IF NOT EXISTS `{databaseName}` " +
            "DEFAULT CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci");
    }

    /// <summary>
    /// 授权业务账号访问新租户库;失败不抛(建表阶段仍可用高权限连接)。
    /// </summary>
    private static void TryGrantTenantDatabaseAccess(
        DbType dbType,
        string privilegedConnectionString,
        string databaseName,
        string tenantConnectionString)
    {
        var appUserId = ExtractConnectionValue(tenantConnectionString, "uid");
        var privilegedUserId = ExtractConnectionValue(privilegedConnectionString, "uid");
        if (string.IsNullOrWhiteSpace(appUserId))
        {
            return;
        }

        if (string.Equals(appUserId, privilegedUserId, StringComparison.OrdinalIgnoreCase))
        {
            return;
        }

        try
        {
            using var grantDb = new SqlSugarClient(new ConnectionConfig
            {
                ConfigId = $"tenant-grant-{databaseName}-{Guid.NewGuid():N}",
                DbType = dbType,
                ConnectionString = privilegedConnectionString,
                IsAutoCloseConnection = true
            });

            grantDb.Ado.ExecuteCommand(
                $"GRANT ALL PRIVILEGES ON `{databaseName}`.* TO '{EscapeSqlLiteral(appUserId)}'@'%'");
        }
        catch
        {
            // GRANT 失败不阻断:后续 Init 仍可用高权限连接建表
        }
    }

    private static string BuildServerLevelConnectionString(string connectionString)
    {
        var segments = connectionString
            .Split(';', StringSplitOptions.RemoveEmptyEntries)
            .Select(x => x.Trim())
            .Where(x => !x.StartsWith("database=", StringComparison.OrdinalIgnoreCase)
                        && !x.StartsWith("initial catalog=", StringComparison.OrdinalIgnoreCase))
            .ToList();
        return string.Join(';', segments) + ";";
    }

    private static string GetRootMessage(Exception ex)
    {
        var current = ex;
        while (current.InnerException != null)
        {
            current = current.InnerException;
        }

        return current.Message;
    }

    private static string EscapeSqlLiteral(string value)
        => value.Replace("'", "''", StringComparison.Ordinal);

    private static string? ExtractConnectionValue(string connectionString, string key)
    {
        if (string.IsNullOrWhiteSpace(connectionString))
        {
            return null;
        }

        var prefix = key + "=";
        foreach (var segment in connectionString.Split(';', StringSplitOptions.RemoveEmptyEntries))
        {
            var part = segment.Trim();
            if (part.StartsWith(prefix, StringComparison.OrdinalIgnoreCase))
            {
                return part[prefix.Length..].Trim();
            }
        }

        return null;
    }
}