feat: add Go migration compatibility foundation
This commit is contained in:
1 parent
7de3300034
commit
716a8031b1
27 files changed
+2582
No files matched your search
@@ -0,0 +1,165 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Backend string
|
||||
|
||||
const (
|
||||
BackendLocal Backend = "local"
|
||||
BackendPostgres Backend = "postgres"
|
||||
)
|
||||
|
||||
type SSLMode string
|
||||
|
||||
const (
|
||||
SSLDisable SSLMode = "disable"
|
||||
SSLVerifyFull SSLMode = "verify-full"
|
||||
)
|
||||
|
||||
type Getenv func(string) string
|
||||
type ReadFile func(string) ([]byte, error)
|
||||
|
||||
type Config struct {
|
||||
Backend Backend
|
||||
DatabaseURL string
|
||||
SSLMode SSLMode
|
||||
TLSConfig *tls.Config
|
||||
PoolMax int32
|
||||
IdleTimeout time.Duration
|
||||
ConnectionTimeout time.Duration
|
||||
StatementTimeout time.Duration
|
||||
ApplicationName string
|
||||
}
|
||||
|
||||
func ParseConfig(getenv Getenv, readFile ReadFile) (Config, error) {
|
||||
if getenv == nil {
|
||||
return Config{}, fmt.Errorf("environment getter is required")
|
||||
}
|
||||
cfg := Config{
|
||||
PoolMax: 10,
|
||||
IdleTimeout: 30 * time.Second,
|
||||
ConnectionTimeout: 10 * time.Second,
|
||||
StatementTimeout: 30 * time.Second,
|
||||
ApplicationName: "zhinian-go",
|
||||
SSLMode: SSLDisable,
|
||||
}
|
||||
|
||||
backend := strings.ToLower(strings.TrimSpace(getenv("ZHINIAN_DATA_BACKEND")))
|
||||
if backend == "" && strings.ToLower(strings.TrimSpace(getenv("NODE_ENV"))) != "production" {
|
||||
backend = string(BackendLocal)
|
||||
}
|
||||
switch Backend(backend) {
|
||||
case BackendLocal, BackendPostgres:
|
||||
cfg.Backend = Backend(backend)
|
||||
default:
|
||||
return Config{}, fmt.Errorf("ZHINIAN_DATA_BACKEND must be explicitly set to 'local' or 'postgres'")
|
||||
}
|
||||
|
||||
if cfg.Backend == BackendLocal {
|
||||
return cfg, nil
|
||||
}
|
||||
var err error
|
||||
if cfg.PoolMax, err = positiveInt32(getenv, "DATABASE_POOL_MAX", cfg.PoolMax); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.IdleTimeout, err = nonNegativeMilliseconds(getenv, "DATABASE_IDLE_TIMEOUT_MS", cfg.IdleTimeout); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.ConnectionTimeout, err = positiveMilliseconds(getenv, "DATABASE_CONNECTION_TIMEOUT_MS", cfg.ConnectionTimeout); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.StatementTimeout, err = positiveMilliseconds(getenv, "DATABASE_STATEMENT_TIMEOUT_MS", cfg.StatementTimeout); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if value := strings.TrimSpace(getenv("DATABASE_APPLICATION_NAME")); value != "" {
|
||||
cfg.ApplicationName = value
|
||||
}
|
||||
cfg.DatabaseURL = strings.TrimSpace(getenv("DATABASE_URL"))
|
||||
if cfg.DatabaseURL == "" {
|
||||
return Config{}, fmt.Errorf("DATABASE_URL is required when ZHINIAN_DATA_BACKEND=postgres")
|
||||
}
|
||||
parsed, parseErr := url.ParseRequestURI(cfg.DatabaseURL)
|
||||
if parseErr != nil || parsed.Host == "" {
|
||||
return Config{}, fmt.Errorf("DATABASE_URL must be a valid PostgreSQL connection URI")
|
||||
}
|
||||
if parsed.Scheme != "postgres" && parsed.Scheme != "postgresql" {
|
||||
return Config{}, fmt.Errorf("DATABASE_URL must use the postgres:// or postgresql:// scheme")
|
||||
}
|
||||
for key := range parsed.Query() {
|
||||
if strings.HasPrefix(strings.ToLower(key), "ssl") {
|
||||
return Config{}, fmt.Errorf("DATABASE_URL must not contain SSL query parameters (%s); use DATABASE_SSL_MODE and DATABASE_CA_CERT_PATH", key)
|
||||
}
|
||||
}
|
||||
|
||||
mode := strings.ToLower(strings.TrimSpace(getenv("DATABASE_SSL_MODE")))
|
||||
if mode != "" {
|
||||
cfg.SSLMode = SSLMode(mode)
|
||||
}
|
||||
switch cfg.SSLMode {
|
||||
case SSLDisable:
|
||||
case SSLVerifyFull:
|
||||
path := strings.TrimSpace(getenv("DATABASE_CA_CERT_PATH"))
|
||||
if path == "" {
|
||||
return Config{}, fmt.Errorf("DATABASE_CA_CERT_PATH is required when DATABASE_SSL_MODE=verify-full")
|
||||
}
|
||||
if readFile == nil {
|
||||
return Config{}, fmt.Errorf("CA certificate reader is required")
|
||||
}
|
||||
pem, readErr := readFile(path)
|
||||
if readErr != nil {
|
||||
return Config{}, fmt.Errorf("read DATABASE_CA_CERT_PATH: %w", readErr)
|
||||
}
|
||||
roots := x509.NewCertPool()
|
||||
if !roots.AppendCertsFromPEM(pem) {
|
||||
return Config{}, fmt.Errorf("DATABASE_CA_CERT_PATH does not contain a valid CA certificate")
|
||||
}
|
||||
cfg.TLSConfig = &tls.Config{RootCAs: roots, MinVersion: tls.VersionTLS12}
|
||||
default:
|
||||
return Config{}, fmt.Errorf("DATABASE_SSL_MODE must be 'disable' or 'verify-full'")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func positiveInt32(getenv Getenv, name string, fallback int32) (int32, error) {
|
||||
raw := strings.TrimSpace(getenv(name))
|
||||
if raw == "" {
|
||||
return fallback, nil
|
||||
}
|
||||
value, err := strconv.ParseInt(raw, 10, 32)
|
||||
if err != nil || value <= 0 {
|
||||
return 0, fmt.Errorf("%s must be a positive integer", name)
|
||||
}
|
||||
return int32(value), nil
|
||||
}
|
||||
|
||||
func nonNegativeMilliseconds(getenv Getenv, name string, fallback time.Duration) (time.Duration, error) {
|
||||
return milliseconds(getenv, name, fallback, true)
|
||||
}
|
||||
|
||||
func positiveMilliseconds(getenv Getenv, name string, fallback time.Duration) (time.Duration, error) {
|
||||
return milliseconds(getenv, name, fallback, false)
|
||||
}
|
||||
|
||||
func milliseconds(getenv Getenv, name string, fallback time.Duration, allowZero bool) (time.Duration, error) {
|
||||
raw := strings.TrimSpace(getenv(name))
|
||||
if raw == "" {
|
||||
return fallback, nil
|
||||
}
|
||||
value, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil || value < 0 || (!allowZero && value == 0) || value > int64((1<<63-1)/time.Millisecond) {
|
||||
qualifier := "positive"
|
||||
if allowZero {
|
||||
qualifier = "non-negative"
|
||||
}
|
||||
return 0, fmt.Errorf("%s must be a %s integer", name, qualifier)
|
||||
}
|
||||
return time.Duration(value) * time.Millisecond, nil
|
||||
}
|
||||
Reference in new issue
Block a user