package database import ( "context" "errors" "fmt" "net/url" "os" "path/filepath" "strconv" "strings" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/stdlib" "gorm.io/driver/postgres" "gorm.io/gorm" ) const ( expectedDatabaseName = "ops_licence" databaseApplicationName = "ops-licence" databaseConnectTimeout = 10 * time.Second ) var pgEnvironmentVariables = [...]string{ "PGHOST", "PGPORT", "PGDATABASE", "PGUSER", "PGPASSWORD", "PGPASSFILE", "PGAPPNAME", "PGCONNECT_TIMEOUT", "PGSSLMODE", "PGSSLKEY", "PGSSLCERT", "PGSSLSNI", "PGSSLROOTCERT", "PGSSLPASSWORD", "PGTARGETSESSIONATTRS", "PGSERVICE", "PGSERVICEFILE", } type databaseTarget struct { normalizedURL string username string password string host string port uint16 sslMode string } func Open(ctx context.Context, dsn string) (*gorm.DB, error) { target, err := parseDatabaseTarget(dsn) if err != nil { return nil, err } if err := rejectPGEnvironmentVariables(); err != nil { return nil, err } config, err := pgx.ParseConfig(target.normalizedURL) if err != nil { return nil, errors.New("解析数据库连接配置失败") } if err := validatePGXConfig(config, target); err != nil { return nil, err } sqlDB := stdlib.OpenDB(*config) closeOnFailure := true defer func() { if closeOnFailure { _ = sqlDB.Close() } }() sqlDB.SetMaxOpenConns(10) sqlDB.SetMaxIdleConns(5) sqlDB.SetConnMaxLifetime(30 * time.Minute) sqlDB.SetConnMaxIdleTime(5 * time.Minute) db, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{ DisableAutomaticPing: true, }) if err != nil { return nil, errors.New("初始化数据库连接失败") } if err := sqlDB.PingContext(ctx); err != nil { return nil, errors.New("连接数据库失败") } var databaseName string if err := sqlDB.QueryRowContext(ctx, "SELECT current_database()").Scan(&databaseName); err != nil { return nil, errors.New("读取当前数据库名称失败") } if databaseName != expectedDatabaseName { return nil, fmt.Errorf("数据库名称必须为 %q,实际为 %q", expectedDatabaseName, databaseName) } closeOnFailure = false return db, nil } func parseDatabaseTarget(dsn string) (databaseTarget, error) { parsed, err := url.ParseRequestURI(dsn) if err != nil || !parsed.IsAbs() || parsed.Opaque != "" { return databaseTarget{}, errors.New("数据库 DSN 必须是完整的 PostgreSQL URL") } if parsed.Scheme != "postgres" && parsed.Scheme != "postgresql" { return databaseTarget{}, errors.New("数据库 DSN 协议必须是 postgres 或 postgresql") } if parsed.Fragment != "" || parsed.RawFragment != "" || strings.Contains(dsn, "#") { return databaseTarget{}, errors.New("数据库 DSN 不能包含 fragment") } if parsed.User == nil || parsed.User.Username() == "" { return databaseTarget{}, errors.New("数据库 DSN 必须显式包含非空用户名") } password, hasPassword := parsed.User.Password() if !hasPassword || password == "" { return databaseTarget{}, errors.New("数据库 DSN 必须显式包含非空密码") } if parsed.Host == "" || strings.Contains(parsed.Host, ",") { return databaseTarget{}, errors.New("数据库 DSN 必须包含单一主机") } host := parsed.Hostname() if host == "" { return databaseTarget{}, errors.New("数据库 DSN 必须显式包含非空主机") } portText := parsed.Port() port, err := strconv.ParseUint(portText, 10, 16) if err != nil || port == 0 { return databaseTarget{}, errors.New("数据库 DSN 必须显式包含有效端口") } if parsed.Path != "/"+expectedDatabaseName || parsed.EscapedPath() != "/"+expectedDatabaseName { return databaseTarget{}, fmt.Errorf("数据库 DSN 的数据库路径必须精确为 /%s", expectedDatabaseName) } query, err := url.ParseQuery(parsed.RawQuery) if err != nil { return databaseTarget{}, errors.New("数据库 DSN 查询参数无效") } for name, values := range query { if !isAllowedDatabaseOption(name) { return databaseTarget{}, errors.New("数据库 DSN 包含不允许的连接选项") } if len(values) != 1 || values[0] == "" { return databaseTarget{}, errors.New("数据库 DSN 的连接选项必须是单个非空值") } } sslModeValues, exists := query["sslmode"] if !exists || !isAllowedSSLMode(sslModeValues[0]) { return databaseTarget{}, errors.New("数据库 DSN 的 sslmode 必须是 disable、require、verify-ca 或 verify-full") } sslMode := sslModeValues[0] rootCertificate := query["sslrootcert"] clientCertificate := query["sslcert"] clientKey := query["sslkey"] sslPassword := query["sslpassword"] for _, option := range []string{"sslrootcert", "sslcert", "sslkey"} { if values, exists := query[option]; exists && !filepath.IsAbs(values[0]) { return databaseTarget{}, errors.New("数据库 DSN 的 TLS 文件路径必须是绝对路径") } } if (len(clientCertificate) == 1) != (len(clientKey) == 1) { return databaseTarget{}, errors.New("数据库 DSN 的 sslcert 与 sslkey 必须成对配置") } if len(sslPassword) == 1 && len(clientKey) == 0 { return databaseTarget{}, errors.New("数据库 DSN 的 sslpassword 必须与 sslkey 一同配置") } if (sslMode == "verify-ca" || sslMode == "verify-full") && len(rootCertificate) == 0 { return databaseTarget{}, errors.New("verify-ca 和 verify-full 必须配置 sslrootcert") } if sslMode == "disable" && (len(rootCertificate) != 0 || len(clientCertificate) != 0 || len(sslPassword) != 0) { return databaseTarget{}, errors.New("sslmode=disable 不能配置 TLS 证书或密码") } query.Set("passfile", os.DevNull) query.Set("application_name", databaseApplicationName) query.Set("connect_timeout", strconv.Itoa(int(databaseConnectTimeout/time.Second))) query.Set("target_session_attrs", "any") query.Set("sslsni", "1") parsed.RawQuery = query.Encode() parsed.ForceQuery = false return databaseTarget{ normalizedURL: parsed.String(), username: parsed.User.Username(), password: password, host: host, port: uint16(port), sslMode: sslMode, }, nil } func isAllowedDatabaseOption(name string) bool { switch name { case "sslmode", "sslrootcert", "sslcert", "sslkey", "sslpassword": return true default: return false } } func isAllowedSSLMode(mode string) bool { switch mode { case "disable", "require", "verify-ca", "verify-full": return true default: return false } } func rejectPGEnvironmentVariables() error { for _, name := range pgEnvironmentVariables { if _, exists := os.LookupEnv(name); exists { return fmt.Errorf("不允许设置数据库环境变量 %s", name) } } return nil } func validatePGXConfig(config *pgx.ConnConfig, target databaseTarget) error { if config.Host != target.host || config.Port != target.port || config.Database != expectedDatabaseName || config.User != target.username || config.Password != target.password || config.ConnectTimeout != databaseConnectTimeout || config.ValidateConnect != nil || len(config.Fallbacks) != 0 { return errors.New("数据库连接配置解析结果不符合预期") } if len(config.RuntimeParams) != 1 || config.RuntimeParams["application_name"] != databaseApplicationName { return errors.New("数据库连接运行参数不符合预期") } tlsRequired := target.sslMode != "disable" if tlsRequired != (config.TLSConfig != nil) { return errors.New("数据库连接 TLS 配置不符合预期") } if target.sslMode == "verify-full" && (config.TLSConfig.ServerName != target.host || config.TLSConfig.InsecureSkipVerify) { return errors.New("数据库连接 verify-full 配置不符合预期") } return nil }