250 lines
7.6 KiB
Go
250 lines
7.6 KiB
Go
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
|
||
}
|