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
|
|||
|
|
}
|