Files
license/internal/database/database.go
2026-07-31 16:59:55 +08:00

250 lines
7.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}