feat: add Go migration compatibility foundation

This commit is contained in:
zn-admin committed 2026-08-13 10:35:52 +08:00
1 parent 7de3300034
commit 716a8031b1
27 files changed
+2582

No files matched your search

+165
View File
@@ -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
}
+152
View File
@@ -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] }
}
+198
View File
@@ -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
}
+183
View File
@@ -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
}
+75
View File
@@ -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)
+62
View File
@@ -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)
}
}