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
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestParseConfigDefaultsToLocalOutsideProduction(t *testing.T) {
|
||||
cfg, err := ParseConfig(env(map[string]string{"NODE_ENV": "development"}), os.ReadFile)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseConfig() error = %v", err)
|
||||
}
|
||||
if cfg.Backend != BackendLocal {
|
||||
t.Fatalf("Backend = %q, want %q", cfg.Backend, BackendLocal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigRequiresExplicitBackendInProduction(t *testing.T) {
|
||||
_, err := ParseConfig(env(map[string]string{"NODE_ENV": "production"}), os.ReadFile)
|
||||
if err == nil || !strings.Contains(err.Error(), "ZHINIAN_DATA_BACKEND") {
|
||||
t.Fatalf("ParseConfig() error = %v, want backend validation error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigRequiresPostgresURL(t *testing.T) {
|
||||
_, err := ParseConfig(env(map[string]string{"ZHINIAN_DATA_BACKEND": "postgres"}), os.ReadFile)
|
||||
if err == nil || !strings.Contains(err.Error(), "DATABASE_URL") {
|
||||
t.Fatalf("ParseConfig() error = %v, want DATABASE_URL error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigAcceptsOnlyPostgresSchemesAndRejectsSSLQueryParameters(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
url string
|
||||
}{
|
||||
{name: "wrong scheme", url: "https://db.example/app"},
|
||||
{name: "sslmode", url: "postgres://db.example/app?sslmode=require"},
|
||||
{name: "mixed case ssl parameter", url: "postgresql://db.example/app?SSLcert=x"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "postgres",
|
||||
"DATABASE_URL": tt.url,
|
||||
}), os.ReadFile)
|
||||
if err == nil {
|
||||
t.Fatal("ParseConfig() error = nil, want validation error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigBuildsVerifyFullTLSFromCA(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
caPath := filepath.Join(dir, "ca.pem")
|
||||
const ca = "-----BEGIN CERTIFICATE-----\nMIIB\n-----END CERTIFICATE-----\n"
|
||||
read := func(path string) ([]byte, error) {
|
||||
if path != caPath {
|
||||
t.Fatalf("read path = %q, want %q", path, caPath)
|
||||
}
|
||||
return []byte(ca), nil
|
||||
}
|
||||
_, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "postgres",
|
||||
"DATABASE_URL": "postgresql://db.example/app",
|
||||
"DATABASE_SSL_MODE": "verify-full",
|
||||
"DATABASE_CA_CERT_PATH": caPath,
|
||||
}), read)
|
||||
if err == nil || !strings.Contains(err.Error(), "CA certificate") {
|
||||
t.Fatalf("ParseConfig() error = %v, want invalid CA certificate error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigVerifyFullBuildsRootsWithoutDisablingVerification(t *testing.T) {
|
||||
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "test CA"},
|
||||
NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour),
|
||||
IsCA: true, BasicConstraintsValid: true, KeyUsage: x509.KeyUsageCertSign,
|
||||
}
|
||||
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ca := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
|
||||
cfg, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "postgres",
|
||||
"DATABASE_URL": "postgresql://db.example/app",
|
||||
"DATABASE_SSL_MODE": "verify-full",
|
||||
"DATABASE_CA_CERT_PATH": "/ca.pem",
|
||||
}), func(string) ([]byte, error) { return ca, nil })
|
||||
if err != nil {
|
||||
t.Fatalf("ParseConfig() error = %v", err)
|
||||
}
|
||||
if cfg.TLSConfig == nil || cfg.TLSConfig.RootCAs == nil {
|
||||
t.Fatal("TLSConfig.RootCAs is nil")
|
||||
}
|
||||
if cfg.TLSConfig.InsecureSkipVerify {
|
||||
t.Fatal("TLSConfig.InsecureSkipVerify = true, want full certificate and hostname verification")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConfigDefaultsAndNumericValidation(t *testing.T) {
|
||||
cfg, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "postgres",
|
||||
"DATABASE_URL": "postgres://db.example/app",
|
||||
}), os.ReadFile)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseConfig() error = %v", err)
|
||||
}
|
||||
if cfg.PoolMax != 10 || cfg.IdleTimeout != 30*time.Second || cfg.ConnectionTimeout != 10*time.Second || cfg.StatementTimeout != 30*time.Second {
|
||||
t.Fatalf("unexpected defaults: %+v", cfg)
|
||||
}
|
||||
if cfg.ApplicationName != "zhinian-go" || cfg.SSLMode != SSLDisable {
|
||||
t.Fatalf("unexpected identity/TLS defaults: %+v", cfg)
|
||||
}
|
||||
|
||||
for name, value := range map[string]string{
|
||||
"DATABASE_POOL_MAX": "0",
|
||||
"DATABASE_IDLE_TIMEOUT_MS": "-1",
|
||||
"DATABASE_CONNECTION_TIMEOUT_MS": "0",
|
||||
"DATABASE_STATEMENT_TIMEOUT_MS": "nope",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "postgres",
|
||||
"DATABASE_URL": "postgres://db.example/app",
|
||||
name: value,
|
||||
}), os.ReadFile)
|
||||
if err == nil || !strings.Contains(err.Error(), name) {
|
||||
t.Fatalf("ParseConfig() error = %v, want %s validation error", err, name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func env(values map[string]string) Getenv {
|
||||
return func(name string) string { return values[name] }
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const ReadinessSQL = `
|
||||
WITH required_table_privileges(table_name, privilege_name) AS (
|
||||
VALUES
|
||||
('assets', 'SELECT'), ('assets', 'INSERT'), ('assets', 'DELETE'),
|
||||
('generation_jobs', 'SELECT'), ('generation_jobs', 'INSERT'), ('generation_jobs', 'UPDATE'), ('generation_jobs', 'DELETE'),
|
||||
('usage_events', 'SELECT'), ('usage_events', 'INSERT'), ('usage_events', 'UPDATE'),
|
||||
('projects', 'SELECT'), ('projects', 'UPDATE'),
|
||||
('image_templates', 'SELECT'), ('image_templates', 'INSERT'), ('image_templates', 'UPDATE'), ('image_templates', 'DELETE'),
|
||||
('platform_organizations', 'SELECT'), ('platform_organizations', 'INSERT'), ('platform_organizations', 'UPDATE'), ('platform_organizations', 'DELETE'),
|
||||
('platform_users', 'SELECT'), ('platform_users', 'INSERT'), ('platform_users', 'UPDATE'), ('platform_users', 'DELETE'),
|
||||
('platform_account_migrations', 'SELECT'), ('platform_account_migrations', 'INSERT'), ('platform_account_migrations', 'UPDATE'),
|
||||
('billing_price_rules', 'SELECT'), ('billing_price_rules', 'INSERT'), ('billing_price_rules', 'UPDATE'),
|
||||
('billing_wallets', 'SELECT'), ('billing_wallets', 'INSERT'), ('billing_wallets', 'UPDATE'),
|
||||
('billing_ledger', 'SELECT'), ('billing_ledger', 'INSERT')
|
||||
)
|
||||
SELECT
|
||||
NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM required_table_privileges
|
||||
WHERE to_regclass('public.' || table_name) IS NULL
|
||||
OR NOT has_table_privilege(current_user, 'public.' || table_name, privilege_name)
|
||||
)
|
||||
AND has_function_privilege(
|
||||
current_user,
|
||||
'public.claim_generation_jobs(text,integer,integer)',
|
||||
'EXECUTE'
|
||||
)
|
||||
AND has_function_privilege(
|
||||
current_user,
|
||||
'public.billing_post_wallet_entry(text,text,text,text,text,bigint,text,text,text,jsonb)',
|
||||
'EXECUTE'
|
||||
) AS ready
|
||||
`
|
||||
|
||||
const ClaimGenerationJobsSQL = `SELECT id FROM public.claim_generation_jobs($1::text, $2::integer, $3::integer)`
|
||||
|
||||
const PostWalletEntrySQL = `SELECT ledger_id, balance_after_fen, balance_fen, total_recharged_fen, total_charged_fen, created_at, updated_at, delta_fen FROM public.billing_post_wallet_entry($1::text, $2::text, $3::text, $4::text, $5::text, $6::bigint, $7::text, $8::text, $9::text, $10::jsonb)`
|
||||
|
||||
type Rows interface {
|
||||
Close()
|
||||
Err() error
|
||||
Next() bool
|
||||
Scan(dest ...any) error
|
||||
}
|
||||
|
||||
type Querier interface {
|
||||
Query(context.Context, string, ...any) (Rows, error)
|
||||
}
|
||||
|
||||
type Pool interface {
|
||||
Querier
|
||||
Close()
|
||||
}
|
||||
|
||||
type Database struct {
|
||||
config Config
|
||||
querier Querier
|
||||
}
|
||||
|
||||
type Store = Database
|
||||
|
||||
func NewDatabase(config Config, querier Querier) *Database {
|
||||
return &Database{config: config, querier: querier}
|
||||
}
|
||||
|
||||
func (db *Database) Readiness(ctx context.Context) error {
|
||||
if db.config.Backend == BackendLocal {
|
||||
return nil
|
||||
}
|
||||
if db.querier == nil {
|
||||
return fmt.Errorf("PostgreSQL pool is not open")
|
||||
}
|
||||
rows, err := db.querier.Query(ctx, ReadinessSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("query PostgreSQL readiness: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return fmt.Errorf("read PostgreSQL readiness: %w", err)
|
||||
}
|
||||
return fmt.Errorf("PostgreSQL readiness query returned no row")
|
||||
}
|
||||
var ready bool
|
||||
if err := rows.Scan(&ready); err != nil {
|
||||
return fmt.Errorf("scan PostgreSQL readiness: %w", err)
|
||||
}
|
||||
if !ready {
|
||||
return fmt.Errorf("PostgreSQL schema or application privileges are not ready")
|
||||
}
|
||||
return rows.Err()
|
||||
}
|
||||
|
||||
type GenerationJob struct {
|
||||
ID string
|
||||
}
|
||||
|
||||
func (db *Database) ClaimGenerationJobs(ctx context.Context, workerID string, limit, lockTimeoutSeconds int) ([]GenerationJob, error) {
|
||||
if db.config.Backend != BackendPostgres || db.querier == nil {
|
||||
return nil, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
||||
}
|
||||
limit = max(1, min(limit, 20))
|
||||
rows, err := db.querier.Query(ctx, ClaimGenerationJobsSQL, workerID, limit, lockTimeoutSeconds)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("claim generation jobs: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var jobs []GenerationJob
|
||||
for rows.Next() {
|
||||
var job GenerationJob
|
||||
if err := rows.Scan(&job.ID); err != nil {
|
||||
return nil, fmt.Errorf("scan claimed generation job: %w", err)
|
||||
}
|
||||
jobs = append(jobs, job)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("read claimed generation jobs: %w", err)
|
||||
}
|
||||
return jobs, nil
|
||||
}
|
||||
|
||||
type WalletEntryParams struct {
|
||||
LedgerID string
|
||||
OrganizationID string
|
||||
AccountID string
|
||||
JobID string
|
||||
Kind string
|
||||
DeltaFen int64
|
||||
Currency string
|
||||
IdempotencyKey string
|
||||
Description string
|
||||
Metadata json.RawMessage
|
||||
}
|
||||
|
||||
type WalletEntry struct {
|
||||
LedgerID string
|
||||
BalanceAfterFen int64
|
||||
BalanceFen int64
|
||||
TotalRechargedFen int64
|
||||
TotalChargedFen int64
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
DeltaFen int64
|
||||
}
|
||||
|
||||
func (db *Database) PostWalletEntry(ctx context.Context, params WalletEntryParams) (WalletEntry, error) {
|
||||
if db.config.Backend != BackendPostgres || db.querier == nil {
|
||||
return WalletEntry{}, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
||||
}
|
||||
accountID := optionalDatabaseText(params.AccountID)
|
||||
if params.Kind == "recharge" || params.Kind == "adjustment" {
|
||||
accountID = nil
|
||||
}
|
||||
jobID := optionalDatabaseText(params.JobID)
|
||||
currency := params.Currency
|
||||
if currency == "" {
|
||||
currency = "CNY"
|
||||
}
|
||||
rows, err := db.querier.Query(ctx, PostWalletEntrySQL,
|
||||
params.LedgerID, params.OrganizationID, accountID, jobID, params.Kind,
|
||||
params.DeltaFen, currency, params.IdempotencyKey, params.Description, params.Metadata,
|
||||
)
|
||||
if err != nil {
|
||||
return WalletEntry{}, fmt.Errorf("post wallet entry: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return WalletEntry{}, fmt.Errorf("read wallet entry: %w", err)
|
||||
}
|
||||
return WalletEntry{}, fmt.Errorf("billing_post_wallet_entry returned no row")
|
||||
}
|
||||
var entry WalletEntry
|
||||
if err := rows.Scan(
|
||||
&entry.LedgerID, &entry.BalanceAfterFen, &entry.BalanceFen,
|
||||
&entry.TotalRechargedFen, &entry.TotalChargedFen, &entry.CreatedAt,
|
||||
&entry.UpdatedAt, &entry.DeltaFen,
|
||||
); err != nil {
|
||||
return WalletEntry{}, fmt.Errorf("scan wallet entry: %w", err)
|
||||
}
|
||||
return entry, rows.Err()
|
||||
}
|
||||
|
||||
func optionalDatabaseText(value string) any {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReadinessSQLFreezesPrivilegeMatrixAndFunctionSignatures(t *testing.T) {
|
||||
wantTables := []string{
|
||||
"assets", "generation_jobs", "usage_events", "projects", "image_templates",
|
||||
"platform_organizations", "platform_users", "platform_account_migrations",
|
||||
"billing_price_rules", "billing_wallets", "billing_ledger",
|
||||
}
|
||||
for _, table := range wantTables {
|
||||
if !strings.Contains(ReadinessSQL, "('"+table+"',") {
|
||||
t.Errorf("ReadinessSQL missing table %q", table)
|
||||
}
|
||||
}
|
||||
if got := strings.Count(ReadinessSQL, "has_function_privilege("); got != 2 {
|
||||
t.Fatalf("has_function_privilege count = %d, want 2", got)
|
||||
}
|
||||
for _, signature := range []string{
|
||||
"public.claim_generation_jobs(text,integer,integer)",
|
||||
"public.billing_post_wallet_entry(text,text,text,text,text,bigint,text,text,text,jsonb)",
|
||||
} {
|
||||
if !strings.Contains(ReadinessSQL, signature) {
|
||||
t.Errorf("ReadinessSQL missing function signature %q", signature)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadinessLocalSucceedsWithoutQuery(t *testing.T) {
|
||||
db := NewDatabase(Config{Backend: BackendLocal}, &fakeQuerier{err: errors.New("must not query")})
|
||||
if err := db.Readiness(context.Background()); err != nil {
|
||||
t.Fatalf("Readiness() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadinessPostgresFailsWhenMatrixIsNotReady(t *testing.T) {
|
||||
q := &fakeQuerier{rows: [][]any{{false}}}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, q)
|
||||
if err := db.Readiness(context.Background()); err == nil {
|
||||
t.Fatal("Readiness() error = nil, want not-ready error")
|
||||
}
|
||||
if q.sql != ReadinessSQL {
|
||||
t.Fatalf("query = %q, want exact ReadinessSQL", q.sql)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClaimGenerationJobsCallsExactFunction(t *testing.T) {
|
||||
const wantSQL = `SELECT id FROM public.claim_generation_jobs($1::text, $2::integer, $3::integer)`
|
||||
if ClaimGenerationJobsSQL != wantSQL {
|
||||
t.Fatalf("ClaimGenerationJobsSQL = %q, want %q", ClaimGenerationJobsSQL, wantSQL)
|
||||
}
|
||||
q := &fakeQuerier{rows: [][]any{{"job-1"}}}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, q)
|
||||
jobs, err := db.ClaimGenerationJobs(context.Background(), "worker-1", 2, 300)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimGenerationJobs() error = %v", err)
|
||||
}
|
||||
if q.sql != ClaimGenerationJobsSQL || !reflect.DeepEqual(q.args, []any{"worker-1", 2, 300}) {
|
||||
t.Fatalf("query = %q args = %#v", q.sql, q.args)
|
||||
}
|
||||
if !reflect.DeepEqual(jobs, []GenerationJob{{ID: "job-1"}}) {
|
||||
t.Fatalf("jobs = %#v", jobs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClaimGenerationJobsBoundsBatchLikeCurrentBackend(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
requested int
|
||||
want int
|
||||
}{
|
||||
{name: "minimum", requested: 0, want: 1},
|
||||
{name: "maximum", requested: 25, want: 20},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
q := &fakeQuerier{}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, q)
|
||||
if _, err := db.ClaimGenerationJobs(context.Background(), "worker-1", test.requested, 300); err != nil {
|
||||
t.Fatalf("ClaimGenerationJobs() error = %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(q.args, []any{"worker-1", test.want, 300}) {
|
||||
t.Fatalf("args = %#v, want bounded limit %d", q.args, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostWalletEntryCallsExactFunction(t *testing.T) {
|
||||
const wantSQL = `SELECT ledger_id, balance_after_fen, balance_fen, total_recharged_fen, total_charged_fen, created_at, updated_at, delta_fen FROM public.billing_post_wallet_entry($1::text, $2::text, $3::text, $4::text, $5::text, $6::bigint, $7::text, $8::text, $9::text, $10::jsonb)`
|
||||
if PostWalletEntrySQL != wantSQL {
|
||||
t.Fatalf("PostWalletEntrySQL = %q, want %q", PostWalletEntrySQL, wantSQL)
|
||||
}
|
||||
q := &fakeQuerier{rows: [][]any{{"ledger-1", int64(120), int64(120), int64(200), int64(80), nil, nil, int64(-80)}}}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, q)
|
||||
metadata := json.RawMessage(`{"source":"test"}`)
|
||||
entry, err := db.PostWalletEntry(context.Background(), WalletEntryParams{
|
||||
LedgerID: "ledger-1", OrganizationID: "org-1", AccountID: "acct-1", JobID: "job-1",
|
||||
Kind: "charge", DeltaFen: -80, Currency: "CNY", IdempotencyKey: "idem-1",
|
||||
Description: "generation", Metadata: metadata,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("PostWalletEntry() error = %v", err)
|
||||
}
|
||||
if q.sql != PostWalletEntrySQL {
|
||||
t.Fatalf("query = %q, want exact PostWalletEntrySQL", q.sql)
|
||||
}
|
||||
wantArgs := []any{"ledger-1", "org-1", "acct-1", "job-1", "charge", int64(-80), "CNY", "idem-1", "generation", metadata}
|
||||
if !reflect.DeepEqual(q.args, wantArgs) {
|
||||
t.Fatalf("args = %#v, want %#v", q.args, wantArgs)
|
||||
}
|
||||
if entry.LedgerID != "ledger-1" || entry.BalanceFen != 120 {
|
||||
t.Fatalf("entry = %#v", entry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostWalletEntryNormalizesOptionalValuesLikeCurrentBackend(t *testing.T) {
|
||||
q := &fakeQuerier{rows: [][]any{{"ledger-1", int64(200), int64(200), int64(200), int64(0), nil, nil, int64(200)}}}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, q)
|
||||
_, err := db.PostWalletEntry(context.Background(), WalletEntryParams{
|
||||
LedgerID: "ledger-1", OrganizationID: "org-1", Kind: "recharge", DeltaFen: 200,
|
||||
IdempotencyKey: "idem-1", Description: "recharge",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("PostWalletEntry() error = %v", err)
|
||||
}
|
||||
wantArgs := []any{"ledger-1", "org-1", nil, nil, "recharge", int64(200), "CNY", "idem-1", "recharge", json.RawMessage(nil)}
|
||||
if !reflect.DeepEqual(q.args, wantArgs) {
|
||||
t.Fatalf("args = %#v, want %#v", q.args, wantArgs)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeQuerier struct {
|
||||
rows [][]any
|
||||
err error
|
||||
sql string
|
||||
args []any
|
||||
}
|
||||
|
||||
func (q *fakeQuerier) Query(_ context.Context, sql string, args ...any) (Rows, error) {
|
||||
q.sql = sql
|
||||
q.args = args
|
||||
return &fakeRows{rows: q.rows, err: q.err}, nil
|
||||
}
|
||||
|
||||
type fakeRows struct {
|
||||
rows [][]any
|
||||
idx int
|
||||
err error
|
||||
}
|
||||
|
||||
func (r *fakeRows) Close() {}
|
||||
func (r *fakeRows) Err() error { return r.err }
|
||||
func (r *fakeRows) Next() bool { return r.idx < len(r.rows) }
|
||||
func (r *fakeRows) Scan(dest ...any) error {
|
||||
if r.idx >= len(r.rows) {
|
||||
return errors.New("scan past end")
|
||||
}
|
||||
row := r.rows[r.idx]
|
||||
r.idx++
|
||||
if len(dest) != len(row) {
|
||||
return errors.New("scan arity mismatch")
|
||||
}
|
||||
for i := range dest {
|
||||
switch target := dest[i].(type) {
|
||||
case *bool:
|
||||
*target = row[i].(bool)
|
||||
case *string:
|
||||
*target = row[i].(string)
|
||||
case *int64:
|
||||
*target = row[i].(int64)
|
||||
default:
|
||||
// nil timestamp fixtures intentionally leave zero values.
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
type Module struct {
|
||||
pool Pool
|
||||
Store *Store
|
||||
}
|
||||
|
||||
func (module *Module) Close() {
|
||||
if module != nil && module.pool != nil {
|
||||
module.pool.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func Open(ctx context.Context, config Config) (*Module, error) {
|
||||
if config.Backend == BackendLocal {
|
||||
return &Module{Store: NewDatabase(config, nil)}, nil
|
||||
}
|
||||
if config.Backend != BackendPostgres {
|
||||
return nil, fmt.Errorf("unsupported data backend %q", config.Backend)
|
||||
}
|
||||
poolConfig, err := pgxpool.ParseConfig(config.DatabaseURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse DATABASE_URL: %w", err)
|
||||
}
|
||||
poolConfig.MaxConns = config.PoolMax
|
||||
poolConfig.MaxConnIdleTime = config.IdleTimeout
|
||||
poolConfig.ConnConfig.ConnectTimeout = config.ConnectionTimeout
|
||||
poolConfig.ConnConfig.RuntimeParams["statement_timeout"] = strconv.FormatInt(config.StatementTimeout.Milliseconds(), 10)
|
||||
poolConfig.ConnConfig.RuntimeParams["application_name"] = config.ApplicationName
|
||||
switch config.SSLMode {
|
||||
case SSLDisable:
|
||||
poolConfig.ConnConfig.TLSConfig = nil
|
||||
poolConfig.ConnConfig.Fallbacks = nil
|
||||
case SSLVerifyFull:
|
||||
if config.TLSConfig == nil {
|
||||
return nil, fmt.Errorf("TLS configuration is required when DATABASE_SSL_MODE=verify-full")
|
||||
}
|
||||
tlsConfig := config.TLSConfig.Clone()
|
||||
if tlsConfig.ServerName == "" {
|
||||
tlsConfig.ServerName = poolConfig.ConnConfig.Host
|
||||
}
|
||||
poolConfig.ConnConfig.TLSConfig = tlsConfig
|
||||
poolConfig.ConnConfig.Fallbacks = nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported DATABASE_SSL_MODE %q", config.SSLMode)
|
||||
}
|
||||
pool, err := pgxpool.NewWithConfig(ctx, poolConfig)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open PostgreSQL pool: %w", err)
|
||||
}
|
||||
adapter := &pgxPoolAdapter{pool: pool}
|
||||
return &Module{pool: adapter, Store: NewDatabase(config, adapter)}, nil
|
||||
}
|
||||
|
||||
type pgxPoolAdapter struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
func (p *pgxPoolAdapter) Query(ctx context.Context, sql string, args ...any) (Rows, error) {
|
||||
return p.pool.Query(ctx, sql, args...)
|
||||
}
|
||||
|
||||
func (p *pgxPoolAdapter) Close() {
|
||||
p.pool.Close()
|
||||
}
|
||||
|
||||
var _ Pool = (*pgxPoolAdapter)(nil)
|
||||
@@ -0,0 +1,62 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestOpenLocalReturnsStoreWithoutPool(t *testing.T) {
|
||||
module, err := Open(context.Background(), Config{Backend: BackendLocal})
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
if module.Store == nil {
|
||||
t.Fatal("Store = nil")
|
||||
}
|
||||
module.Close()
|
||||
if err := module.Store.Readiness(context.Background()); err != nil {
|
||||
t.Fatalf("Readiness() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenPostgresDoesNotProbeBeforeReadiness(t *testing.T) {
|
||||
module, err := Open(context.Background(), Config{
|
||||
Backend: BackendPostgres,
|
||||
DatabaseURL: "postgresql://app:secret@127.0.0.1:1/app",
|
||||
SSLMode: SSLDisable,
|
||||
PoolMax: 1,
|
||||
ApplicationName: "zhinian-go-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v; connection failures belong to readiness", err)
|
||||
}
|
||||
module.Close()
|
||||
}
|
||||
|
||||
func TestParseConfigIgnoresPostgresTLSSettingsForLocalBackend(t *testing.T) {
|
||||
cfg, err := ParseConfig(env(map[string]string{
|
||||
"ZHINIAN_DATA_BACKEND": "local",
|
||||
"DATABASE_SSL_MODE": "verify-full",
|
||||
"DATABASE_CA_CERT_PATH": "/missing/ca.pem",
|
||||
"DATABASE_POOL_MAX": "0",
|
||||
"DATABASE_CONNECTION_TIMEOUT_MS": "invalid",
|
||||
}), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseConfig() error = %v", err)
|
||||
}
|
||||
if cfg.Backend != BackendLocal || cfg.TLSConfig != nil {
|
||||
t.Fatalf("config = %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenPostgresRejectsInvalidConfiguredTLSMode(t *testing.T) {
|
||||
_, err := Open(context.Background(), Config{
|
||||
Backend: BackendPostgres,
|
||||
DatabaseURL: "postgresql://app:secret@db.example/app",
|
||||
SSLMode: SSLMode("prefer"),
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "DATABASE_SSL_MODE") {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user