mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 23:06:36 +08:00
772 lines
20 KiB
Go
772 lines
20 KiB
Go
// Copyright 2026 Arctel.net
|
|
// SPDX-License-Identifier: Apache-2.0
|
|
|
|
package cmd
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"net/url"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"regexp"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"Wavelet/core"
|
|
|
|
"github.com/glebarez/sqlite"
|
|
"gorm.io/driver/postgres"
|
|
"gorm.io/gorm"
|
|
gormlogger "gorm.io/gorm/logger"
|
|
)
|
|
|
|
const (
|
|
goldenRoot = "/Users/ryan/Code/Go/OpenFlare"
|
|
goldCommit = "9f79fb99"
|
|
goldGooseVersion = int64(202608090003)
|
|
sampleZoneDomain = "l3-upgrade-golden.example"
|
|
goldMigrateWait = 75 * time.Second
|
|
legacyPluginStamp = "openflare/legacy"
|
|
serverPluginStamp = "server"
|
|
)
|
|
|
|
var (
|
|
goldBinOnce sync.Once
|
|
goldBinPath string
|
|
goldSrcDir string
|
|
goldBinErr error
|
|
)
|
|
|
|
func TestUpgradeFromGolden(t *testing.T) {
|
|
t.Run("sqlite", func(t *testing.T) {
|
|
tmp := t.TempDir()
|
|
dbPath := filepath.Join(tmp, "a.db")
|
|
runGoldenAPI(t, tmp, goldSQLiteEnv(t, tmp, dbPath), func() bool {
|
|
return sqliteReady(dbPath)
|
|
})
|
|
assertUpgradeFromGolden(t, upgradeDB{
|
|
sqlitePath: dbPath,
|
|
source: cordisSQLiteSource(t, dbPath),
|
|
})
|
|
})
|
|
}
|
|
|
|
func TestUpgradePostgresFromGolden(t *testing.T) {
|
|
dsn := strings.TrimSpace(os.Getenv("TEST_PG_DSN"))
|
|
if dsn == "" {
|
|
t.Skip("TEST_PG_DSN is not set")
|
|
}
|
|
|
|
host, port, user, pass, adminDB, sslMode := parsePostgresDSN(t, dsn)
|
|
adminDSN := postgresDSN(host, port, user, pass, adminDB, sslMode)
|
|
admin := openInspectDB(t, "", adminDSN)
|
|
t.Cleanup(func() { _ = admin.Close() })
|
|
|
|
dbName := fmt.Sprintf("of_l3_%d", time.Now().UnixNano())
|
|
if !safePGIdent(dbName) {
|
|
t.Fatalf("generated database name %q is not a safe identifier", dbName)
|
|
}
|
|
if _, err := admin.Exec("CREATE DATABASE " + dbName); err != nil {
|
|
t.Fatalf("CREATE DATABASE %s: %v", dbName, err)
|
|
}
|
|
t.Cleanup(func() {
|
|
_, _ = admin.Exec(`SELECT pg_terminate_backend(pid) FROM pg_stat_activity WHERE datname = $1 AND pid <> pg_backend_pid()`, dbName)
|
|
_, _ = admin.Exec("DROP DATABASE IF EXISTS " + dbName)
|
|
})
|
|
|
|
tmp := t.TempDir()
|
|
testDSN := postgresDSN(host, port, user, pass, dbName, sslMode)
|
|
runGoldenAPI(t, tmp, goldPostgresEnv(t, tmp, host, port, user, pass, dbName, sslMode), func() bool {
|
|
return postgresReady(testDSN)
|
|
})
|
|
assertUpgradeFromGolden(t, upgradeDB{
|
|
pgDSN: testDSN,
|
|
source: cordisPostgresSource(t, host, port, user, pass, dbName, sslMode),
|
|
})
|
|
}
|
|
|
|
func TestUpgradePostgresFromExistingDump(t *testing.T) {
|
|
dsn := strings.TrimSpace(os.Getenv("TEST_PG_EXISTING_DSN"))
|
|
if dsn == "" {
|
|
t.Skip("TEST_PG_EXISTING_DSN is not set")
|
|
}
|
|
|
|
host, port, user, pass, dbName, sslMode := parsePostgresDSN(t, dsn)
|
|
spec := upgradeDB{
|
|
pgDSN: dsn,
|
|
source: cordisPostgresSource(t, host, port, user, pass, dbName, sslMode),
|
|
}
|
|
|
|
inspect := openInspectDB(t, "", spec.pgDSN)
|
|
beforeCounts := countNamedTables(t, inspect, productionCountTables)
|
|
beforeTables := listPublicTables(t, inspect)
|
|
_ = inspect.Close()
|
|
|
|
assertUpgradeFromGolden(t, spec)
|
|
|
|
inspect = openInspectDB(t, "", spec.pgDSN)
|
|
defer func() { _ = inspect.Close() }()
|
|
afterCounts := countNamedTables(t, inspect, productionCountTables)
|
|
for _, name := range productionCountTables {
|
|
if afterCounts[name] < beforeCounts[name] {
|
|
t.Errorf("row count dropped for %s: before %d after %d", name, beforeCounts[name], afterCounts[name])
|
|
}
|
|
}
|
|
afterTables := listPublicTables(t, inspect)
|
|
for name := range beforeTables {
|
|
if !afterTables[name] {
|
|
t.Errorf("table %s dropped", name)
|
|
}
|
|
}
|
|
for _, name := range []string{"w_schema_versions", "w_message_channels", "w_message_bindings", "w_message_pairing_codes"} {
|
|
if !afterTables[name] {
|
|
t.Errorf("expected upgrade to create %s", name)
|
|
}
|
|
}
|
|
var n int
|
|
if err := inspect.QueryRow(`SELECT COUNT(*) FROM pg_inherits i JOIN pg_class c ON c.oid = i.inhparent WHERE c.relname IN ('of_node_access_logs', 'w_user_access_logs')`).Scan(&n); err != nil {
|
|
t.Fatalf("count partitions: %v", err)
|
|
}
|
|
if n < 8 {
|
|
t.Errorf("partition children = %d, want at least 8", n)
|
|
}
|
|
}
|
|
|
|
var productionCountTables = []string{
|
|
"of_zones", "of_zone_domains", "of_proxy_routes", "of_nodes", "of_origins",
|
|
"of_tls_certificates", "of_waf_rule_groups", "of_pages_projects",
|
|
"w_users", "w_schedules", "w_system_configs", "w_templates", "w_uploads",
|
|
"of_node_access_logs", "w_user_access_logs",
|
|
}
|
|
|
|
func countNamedTables(t *testing.T, db *sql.DB, tables []string) map[string]int {
|
|
t.Helper()
|
|
out := make(map[string]int, len(tables))
|
|
for _, name := range tables {
|
|
if !safePGIdent(name) {
|
|
t.Fatalf("unsafe table name %q", name)
|
|
}
|
|
var n int
|
|
if err := db.QueryRow("SELECT COUNT(*) FROM " + name).Scan(&n); err != nil {
|
|
t.Fatalf("count %s: %v", name, err)
|
|
}
|
|
out[name] = n
|
|
}
|
|
return out
|
|
}
|
|
|
|
func listPublicTables(t *testing.T, db *sql.DB) map[string]bool {
|
|
t.Helper()
|
|
rows, err := db.Query(`SELECT tablename FROM pg_tables WHERE schemaname = 'public'`)
|
|
if err != nil {
|
|
t.Fatalf("list public tables: %v", err)
|
|
}
|
|
defer func() { _ = rows.Close() }()
|
|
out := make(map[string]bool)
|
|
for rows.Next() {
|
|
var name string
|
|
if err := rows.Scan(&name); err != nil {
|
|
t.Fatalf("scan table name: %v", err)
|
|
}
|
|
out[name] = true
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
t.Fatalf("list public tables: %v", err)
|
|
}
|
|
return out
|
|
}
|
|
|
|
type upgradeDB struct {
|
|
sqlitePath string
|
|
pgDSN string
|
|
source core.ConfigSource
|
|
}
|
|
|
|
func assertUpgradeFromGolden(t *testing.T, spec upgradeDB) {
|
|
t.Helper()
|
|
|
|
inspect := openInspectDB(t, spec.sqlitePath, spec.pgDSN)
|
|
before := dumpOfSchema(t, inspect, spec.pgDSN != "")
|
|
insertSQL := `INSERT INTO of_zones (domain) VALUES (?)`
|
|
if spec.pgDSN != "" {
|
|
insertSQL = `INSERT INTO of_zones (domain) VALUES ($1)`
|
|
}
|
|
if _, err := inspect.Exec(insertSQL, sampleZoneDomain); err != nil {
|
|
t.Fatalf("insert sample of_zones row: %v", err)
|
|
}
|
|
_ = inspect.Close()
|
|
|
|
app := cordisPrepare(t, spec.source)
|
|
legacyRows := schemaPluginRows(t, spec, legacyPluginStamp)
|
|
assertStampedUpgrade(t, spec, before, legacyRows)
|
|
if err := app.Context().Dispose(); err != nil {
|
|
t.Fatalf("dispose first app: %v", err)
|
|
}
|
|
|
|
app2 := cordisPrepare(t, spec.source)
|
|
t.Cleanup(func() { _ = app2.Context().Dispose() })
|
|
if got := schemaPluginRows(t, spec, legacyPluginStamp); got != legacyRows {
|
|
t.Fatalf("second Prepare increased %s rows: got %d, want %d", legacyPluginStamp, got, legacyRows)
|
|
}
|
|
assertStampedUpgrade(t, spec, before, legacyRows)
|
|
}
|
|
|
|
func cordisPrepare(t *testing.T, src core.ConfigSource) *core.App {
|
|
t.Helper()
|
|
app := newOpenFlareApp(core.ProfileAPI, core.WithConfigSource(src))
|
|
if err := app.Prepare(); err != nil {
|
|
t.Fatalf("Prepare: %v", err)
|
|
}
|
|
if err := app.ApplyPlugins(); err != nil {
|
|
t.Fatalf("ApplyPlugins: %v", err)
|
|
}
|
|
if err := app.RunMigrations(); err != nil {
|
|
t.Fatalf("RunMigrations: %v", err)
|
|
}
|
|
return app
|
|
}
|
|
|
|
func assertStampedUpgrade(t *testing.T, spec upgradeDB, before map[string][]string, legacyRows int) {
|
|
t.Helper()
|
|
db := openInspectDB(t, spec.sqlitePath, spec.pgDSN)
|
|
defer func() { _ = db.Close() }()
|
|
postgres := spec.pgDSN != ""
|
|
|
|
if got := gooseMaxVersion(t, db); got != goldGooseVersion {
|
|
t.Errorf("goose_db_version max = %d, want %d", got, goldGooseVersion)
|
|
}
|
|
if legacyRows < 2 {
|
|
t.Errorf("w_schema_versions %s rows = %d, want at least 2 (0 and %d)", legacyPluginStamp, legacyRows, goldGooseVersion)
|
|
}
|
|
if !pluginHasVersion(t, db, postgres, legacyPluginStamp, 0) {
|
|
t.Errorf("missing w_schema_versions (%s, 0)", legacyPluginStamp)
|
|
}
|
|
if !pluginHasVersion(t, db, postgres, legacyPluginStamp, goldGooseVersion) {
|
|
t.Errorf("missing w_schema_versions (%s, %d)", legacyPluginStamp, goldGooseVersion)
|
|
}
|
|
if !pluginHasVersion(t, db, postgres, serverPluginStamp, 1) {
|
|
t.Errorf("missing w_schema_versions (%s, 1)", serverPluginStamp)
|
|
}
|
|
|
|
var domain string
|
|
q := `SELECT domain FROM of_zones WHERE domain = ?`
|
|
if postgres {
|
|
q = `SELECT domain FROM of_zones WHERE domain = $1`
|
|
}
|
|
if err := db.QueryRow(q, sampleZoneDomain).Scan(&domain); err != nil {
|
|
t.Errorf("sample of_zones row missing after upgrade: %v", err)
|
|
}
|
|
|
|
after := dumpOfSchema(t, db, postgres)
|
|
for table, cols := range before {
|
|
got, ok := after[table]
|
|
if !ok {
|
|
t.Errorf("of_* table %s dropped", table)
|
|
continue
|
|
}
|
|
have := make(map[string]bool, len(got))
|
|
for _, c := range got {
|
|
have[c] = true
|
|
}
|
|
for _, c := range cols {
|
|
if !have[c] {
|
|
t.Errorf("of_* column %s.%s dropped", table, c)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func runGoldenAPI(t *testing.T, workDir string, env []string, ready func() bool) {
|
|
t.Helper()
|
|
bin := buildGoldenBinary(t)
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), goldMigrateWait)
|
|
defer cancel()
|
|
cmd := exec.CommandContext(ctx, bin, "api")
|
|
cmd.Dir = workDir
|
|
cmd.Env = env
|
|
var out bytes.Buffer
|
|
cmd.Stdout = &out
|
|
cmd.Stderr = &out
|
|
if err := cmd.Start(); err != nil {
|
|
t.Fatalf("start golden api: %v", err)
|
|
}
|
|
|
|
waitErr := make(chan error, 1)
|
|
go func() { waitErr <- cmd.Wait() }()
|
|
|
|
ticker := time.NewTicker(200 * time.Millisecond)
|
|
defer ticker.Stop()
|
|
for {
|
|
if ready() {
|
|
killGolden(cmd)
|
|
<-waitErr
|
|
return
|
|
}
|
|
select {
|
|
case err := <-waitErr:
|
|
if ready() {
|
|
return
|
|
}
|
|
t.Fatalf("golden api exited before goose %d: %v\n%s", goldGooseVersion, err, out.String())
|
|
case <-ctx.Done():
|
|
killGolden(cmd)
|
|
<-waitErr
|
|
t.Fatalf("timeout waiting for golden goose %d\n%s", goldGooseVersion, out.String())
|
|
case <-ticker.C:
|
|
}
|
|
}
|
|
}
|
|
|
|
func killGolden(cmd *exec.Cmd) {
|
|
if cmd.Process == nil {
|
|
return
|
|
}
|
|
_ = cmd.Process.Kill()
|
|
}
|
|
|
|
func buildGoldenBinary(t *testing.T) string {
|
|
t.Helper()
|
|
goldBinOnce.Do(func() {
|
|
src, err := os.MkdirTemp("", "of-gold-src-")
|
|
if err != nil {
|
|
goldBinErr = err
|
|
return
|
|
}
|
|
archive := exec.Command("git", "-C", goldenRoot, "archive", goldCommit)
|
|
extract := exec.Command("tar", "-x", "-C", src)
|
|
pipe, err := archive.StdoutPipe()
|
|
if err != nil {
|
|
goldBinErr = fmt.Errorf("gold archive pipe: %w", err)
|
|
return
|
|
}
|
|
extract.Stdin = pipe
|
|
var archiveErr, extractErr bytes.Buffer
|
|
archive.Stderr = &archiveErr
|
|
extract.Stderr = &extractErr
|
|
if err := archive.Start(); err != nil {
|
|
goldBinErr = fmt.Errorf("git archive %s: %w", goldCommit, err)
|
|
return
|
|
}
|
|
if err := extract.Start(); err != nil {
|
|
_ = archive.Process.Kill()
|
|
goldBinErr = fmt.Errorf("extract gold %s: %w", goldCommit, err)
|
|
return
|
|
}
|
|
if err := extract.Wait(); err != nil {
|
|
_ = archive.Wait()
|
|
goldBinErr = fmt.Errorf("extract gold %s: %w\n%s", goldCommit, err, extractErr.String())
|
|
return
|
|
}
|
|
if err := archive.Wait(); err != nil {
|
|
goldBinErr = fmt.Errorf("git archive %s: %w\n%s", goldCommit, err, archiveErr.String())
|
|
return
|
|
}
|
|
if _, err := os.Stat(filepath.Join(src, "main.go")); err != nil {
|
|
goldBinErr = fmt.Errorf("gold %s at %s: %w", goldCommit, src, err)
|
|
return
|
|
}
|
|
goldSrcDir = src
|
|
|
|
dir, err := os.MkdirTemp("", "of-gold-bin-")
|
|
if err != nil {
|
|
goldBinErr = err
|
|
return
|
|
}
|
|
out := filepath.Join(dir, "gold")
|
|
cmd := exec.Command("go", "build", "-o", out, ".")
|
|
cmd.Dir = src
|
|
var buf bytes.Buffer
|
|
cmd.Stdout = &buf
|
|
cmd.Stderr = &buf
|
|
if err := cmd.Run(); err != nil {
|
|
goldBinErr = fmt.Errorf("go build golden %s: %w\n%s", goldCommit, err, buf.String())
|
|
return
|
|
}
|
|
goldBinPath = out
|
|
})
|
|
if goldBinErr != nil {
|
|
t.Fatalf("%v", goldBinErr)
|
|
}
|
|
return goldBinPath
|
|
}
|
|
|
|
func copyGoldConfig(t *testing.T, dir string) string {
|
|
t.Helper()
|
|
buildGoldenBinary(t)
|
|
dst := filepath.Join(dir, "config.yaml")
|
|
src, err := os.Open(filepath.Join(goldSrcDir, "config.example.yaml")) //nolint:gosec // extracted gold snapshot
|
|
if err != nil {
|
|
t.Fatalf("open golden config.example.yaml: %v", err)
|
|
}
|
|
defer func() { _ = src.Close() }()
|
|
out, err := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) //nolint:gosec // test temp file
|
|
if err != nil {
|
|
t.Fatalf("create temp config.yaml: %v", err)
|
|
}
|
|
if _, err := io.Copy(out, src); err != nil {
|
|
_ = out.Close()
|
|
t.Fatalf("copy golden config: %v", err)
|
|
}
|
|
if err := out.Close(); err != nil {
|
|
t.Fatalf("close temp config.yaml: %v", err)
|
|
}
|
|
return dst
|
|
}
|
|
|
|
func goldSQLiteEnv(t *testing.T, dir, dbPath string) []string {
|
|
t.Helper()
|
|
cfg := copyGoldConfig(t, dir)
|
|
addr := freeLocalAddr(t)
|
|
return filteredGoldEnv(
|
|
"CONFIG_PATH="+cfg,
|
|
"SQLITE_PATH="+dbPath,
|
|
"DB_ENABLED=false",
|
|
"REDIS_ENABLED=false",
|
|
"CLICKHOUSE_ENABLED=false",
|
|
"APP_ENV=testing",
|
|
"APP_ADDR="+addr,
|
|
)
|
|
}
|
|
|
|
func goldPostgresEnv(t *testing.T, dir, host string, port int, user, pass, dbName, sslMode string) []string {
|
|
t.Helper()
|
|
cfg := copyGoldConfig(t, dir)
|
|
addr := freeLocalAddr(t)
|
|
return filteredGoldEnv(
|
|
"CONFIG_PATH="+cfg,
|
|
"DB_ENABLED=true",
|
|
"DB_HOST="+host,
|
|
"DB_PORT="+strconv.Itoa(port),
|
|
"DB_USERNAME="+user,
|
|
"DB_PASSWORD="+pass,
|
|
"DB_NAME="+dbName,
|
|
"DB_SSL_MODE="+sslMode,
|
|
"REDIS_ENABLED=false",
|
|
"CLICKHOUSE_ENABLED=false",
|
|
"APP_ENV=testing",
|
|
"APP_ADDR="+addr,
|
|
)
|
|
}
|
|
|
|
func filteredGoldEnv(extra ...string) []string {
|
|
drop := map[string]bool{
|
|
"CONFIG_PATH": true,
|
|
"SQLITE_PATH": true,
|
|
"DB_ENABLED": true,
|
|
"DB_HOST": true,
|
|
"DB_PORT": true,
|
|
"DB_USERNAME": true,
|
|
"DB_PASSWORD": true,
|
|
"DB_NAME": true,
|
|
"DB_SSL_MODE": true,
|
|
"REDIS_ENABLED": true,
|
|
"REDIS_ADDR": true,
|
|
"CLICKHOUSE_ENABLED": true,
|
|
"CLICKHOUSE_HOST": true,
|
|
"APP_ENV": true,
|
|
"APP_ADDR": true,
|
|
}
|
|
env := make([]string, 0, len(os.Environ())+len(extra))
|
|
for _, kv := range os.Environ() {
|
|
k, _, _ := strings.Cut(kv, "=")
|
|
if drop[k] {
|
|
continue
|
|
}
|
|
env = append(env, kv)
|
|
}
|
|
return append(env, extra...)
|
|
}
|
|
|
|
func freeLocalAddr(t *testing.T) string {
|
|
t.Helper()
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatalf("listen for free port: %v", err)
|
|
}
|
|
addr := ln.Addr().String()
|
|
_ = ln.Close()
|
|
return addr
|
|
}
|
|
|
|
func cordisSQLiteSource(t *testing.T, dbPath string) core.ConfigSource {
|
|
t.Helper()
|
|
return core.NewMapSource(map[string]any{
|
|
"app": map[string]any{
|
|
"addr": "127.0.0.1:0",
|
|
"env": "testing",
|
|
},
|
|
"redis": map[string]any{
|
|
"enabled": false,
|
|
},
|
|
"clickhouse": map[string]any{
|
|
"enabled": false,
|
|
},
|
|
"database": map[string]any{
|
|
"enabled": false,
|
|
"sqlite_path": dbPath,
|
|
},
|
|
})
|
|
}
|
|
|
|
func cordisPostgresSource(t *testing.T, host string, port int, user, pass, dbName, sslMode string) core.ConfigSource {
|
|
t.Helper()
|
|
return core.NewMapSource(map[string]any{
|
|
"app": map[string]any{
|
|
"addr": "127.0.0.1:0",
|
|
"env": "testing",
|
|
},
|
|
"redis": map[string]any{
|
|
"enabled": false,
|
|
},
|
|
"clickhouse": map[string]any{
|
|
"enabled": false,
|
|
},
|
|
"database": map[string]any{
|
|
"enabled": true,
|
|
"host": host,
|
|
"port": port,
|
|
"username": user,
|
|
"password": pass,
|
|
"database": dbName,
|
|
"ssl_mode": sslMode,
|
|
},
|
|
})
|
|
}
|
|
|
|
func sqliteReady(path string) bool {
|
|
if _, err := os.Stat(path); err != nil {
|
|
return false
|
|
}
|
|
gdb, err := gorm.Open(sqlite.Open("file:"+path+"?mode=ro&_pragma=busy_timeout(1000)"), &gorm.Config{Logger: gormlogger.Default.LogMode(gormlogger.Silent)})
|
|
if err != nil {
|
|
return false
|
|
}
|
|
sqlDB, err := gdb.DB()
|
|
if err != nil {
|
|
return false
|
|
}
|
|
defer func() { _ = sqlDB.Close() }()
|
|
return migratedReady(sqlDB, false)
|
|
}
|
|
|
|
func postgresReady(dsn string) bool {
|
|
gdb, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: gormlogger.Default.LogMode(gormlogger.Silent)})
|
|
if err != nil {
|
|
return false
|
|
}
|
|
sqlDB, err := gdb.DB()
|
|
if err != nil {
|
|
return false
|
|
}
|
|
defer func() { _ = sqlDB.Close() }()
|
|
return migratedReady(sqlDB, true)
|
|
}
|
|
|
|
func migratedReady(db *sql.DB, postgres bool) bool {
|
|
if gooseMaxVersionSilent(db) != goldGooseVersion {
|
|
return false
|
|
}
|
|
var n int
|
|
var err error
|
|
if postgres {
|
|
err = db.QueryRow(`SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = 'public' AND table_name = 'of_nodes'`).Scan(&n)
|
|
} else {
|
|
err = db.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'of_nodes'`).Scan(&n)
|
|
}
|
|
return err == nil && n > 0
|
|
}
|
|
|
|
func openInspectDB(t *testing.T, sqlitePath, pgDSN string) *sql.DB {
|
|
t.Helper()
|
|
var gdb *gorm.DB
|
|
var err error
|
|
if pgDSN != "" {
|
|
gdb, err = gorm.Open(postgres.Open(pgDSN), &gorm.Config{Logger: gormlogger.Default.LogMode(gormlogger.Silent)})
|
|
} else {
|
|
gdb, err = gorm.Open(sqlite.Open(sqlitePath), &gorm.Config{Logger: gormlogger.Default.LogMode(gormlogger.Silent)})
|
|
}
|
|
if err != nil {
|
|
t.Fatalf("open inspect db: %v", err)
|
|
}
|
|
sqlDB, err := gdb.DB()
|
|
if err != nil {
|
|
t.Fatalf("inspect sql.DB: %v", err)
|
|
}
|
|
return sqlDB
|
|
}
|
|
|
|
func dumpOfSchema(t *testing.T, db *sql.DB, postgres bool) map[string][]string {
|
|
t.Helper()
|
|
tables := ofTables(t, db, postgres)
|
|
out := make(map[string][]string, len(tables))
|
|
for _, table := range tables {
|
|
out[table] = ofColumns(t, db, postgres, table)
|
|
}
|
|
if len(out) == 0 {
|
|
t.Fatal("no of_* tables in golden database")
|
|
}
|
|
return out
|
|
}
|
|
|
|
func ofTables(t *testing.T, db *sql.DB, postgres bool) []string {
|
|
t.Helper()
|
|
var rows *sql.Rows
|
|
var err error
|
|
if postgres {
|
|
rows, err = db.Query(`SELECT tablename FROM pg_tables WHERE schemaname = 'public' AND tablename LIKE 'of_%' ORDER BY tablename`)
|
|
} else {
|
|
rows, err = db.Query(`SELECT name FROM sqlite_master WHERE type = 'table' AND name LIKE 'of_%' ORDER BY name`)
|
|
}
|
|
if err != nil {
|
|
t.Fatalf("list of_* tables: %v", err)
|
|
}
|
|
defer func() { _ = rows.Close() }()
|
|
var tables []string
|
|
for rows.Next() {
|
|
var name string
|
|
if err := rows.Scan(&name); err != nil {
|
|
t.Fatalf("scan of_* table: %v", err)
|
|
}
|
|
tables = append(tables, name)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
t.Fatalf("list of_* tables: %v", err)
|
|
}
|
|
return tables
|
|
}
|
|
|
|
func ofColumns(t *testing.T, db *sql.DB, postgres bool, table string) []string {
|
|
t.Helper()
|
|
var rows *sql.Rows
|
|
var err error
|
|
if postgres {
|
|
rows, err = db.Query(`SELECT column_name FROM information_schema.columns WHERE table_schema = 'public' AND table_name = $1 ORDER BY ordinal_position`, table)
|
|
} else {
|
|
rows, err = db.Query(`SELECT name FROM pragma_table_info(?)`, table)
|
|
}
|
|
if err != nil {
|
|
t.Fatalf("list columns for %s: %v", table, err)
|
|
}
|
|
defer func() { _ = rows.Close() }()
|
|
var cols []string
|
|
for rows.Next() {
|
|
var name string
|
|
if err := rows.Scan(&name); err != nil {
|
|
t.Fatalf("scan column for %s: %v", table, err)
|
|
}
|
|
cols = append(cols, name)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
t.Fatalf("list columns for %s: %v", table, err)
|
|
}
|
|
return cols
|
|
}
|
|
|
|
func gooseMaxVersion(t *testing.T, db *sql.DB) int64 {
|
|
t.Helper()
|
|
v := gooseMaxVersionSilent(db)
|
|
if v < 0 {
|
|
t.Fatal("read goose_db_version max failed")
|
|
}
|
|
return v
|
|
}
|
|
|
|
func gooseMaxVersionSilent(db *sql.DB) int64 {
|
|
var v int64
|
|
if err := db.QueryRow(`SELECT COALESCE(MAX(version_id), 0) FROM goose_db_version`).Scan(&v); err != nil {
|
|
return -1
|
|
}
|
|
return v
|
|
}
|
|
|
|
func schemaPluginRows(t *testing.T, spec upgradeDB, pluginID string) int {
|
|
t.Helper()
|
|
db := openInspectDB(t, spec.sqlitePath, spec.pgDSN)
|
|
defer func() { _ = db.Close() }()
|
|
q := `SELECT COUNT(*) FROM w_schema_versions WHERE plugin_id = ?`
|
|
if spec.pgDSN != "" {
|
|
q = `SELECT COUNT(*) FROM w_schema_versions WHERE plugin_id = $1`
|
|
}
|
|
var n int
|
|
if err := db.QueryRow(q, pluginID).Scan(&n); err != nil {
|
|
t.Fatalf("count w_schema_versions %s: %v", pluginID, err)
|
|
}
|
|
return n
|
|
}
|
|
|
|
func pluginHasVersion(t *testing.T, db *sql.DB, postgres bool, pluginID string, version int64) bool {
|
|
t.Helper()
|
|
q := `SELECT COUNT(*) FROM w_schema_versions WHERE plugin_id = ? AND version_id = ?`
|
|
if postgres {
|
|
q = `SELECT COUNT(*) FROM w_schema_versions WHERE plugin_id = $1 AND version_id = $2`
|
|
}
|
|
var n int
|
|
if err := db.QueryRow(q, pluginID, version).Scan(&n); err != nil {
|
|
t.Fatalf("lookup w_schema_versions (%s, %d): %v", pluginID, version, err)
|
|
}
|
|
return n > 0
|
|
}
|
|
|
|
func parsePostgresDSN(t *testing.T, dsn string) (host string, port int, user, pass, dbName, sslMode string) {
|
|
t.Helper()
|
|
u, err := url.Parse(dsn)
|
|
if err != nil {
|
|
t.Fatalf("TEST_PG_DSN: %v", err)
|
|
}
|
|
host = u.Hostname()
|
|
if host == "" {
|
|
host = "127.0.0.1"
|
|
}
|
|
port = 5432
|
|
if p := u.Port(); p != "" {
|
|
port, err = strconv.Atoi(p)
|
|
if err != nil {
|
|
t.Fatalf("TEST_PG_DSN port: %v", err)
|
|
}
|
|
}
|
|
if u.User != nil {
|
|
user = u.User.Username()
|
|
pass, _ = u.User.Password()
|
|
}
|
|
dbName = strings.Trim(u.Path, "/")
|
|
if dbName == "" {
|
|
dbName = "postgres"
|
|
}
|
|
sslMode = u.Query().Get("sslmode")
|
|
if sslMode == "" {
|
|
sslMode = "disable"
|
|
}
|
|
return
|
|
}
|
|
|
|
func postgresDSN(host string, port int, user, pass, dbName, sslMode string) string {
|
|
u := &url.URL{
|
|
Scheme: "postgres",
|
|
Host: net.JoinHostPort(host, strconv.Itoa(port)),
|
|
Path: dbName,
|
|
}
|
|
if user != "" {
|
|
u.User = url.UserPassword(user, pass)
|
|
}
|
|
q := url.Values{}
|
|
q.Set("sslmode", sslMode)
|
|
u.RawQuery = q.Encode()
|
|
return u.String()
|
|
}
|
|
|
|
var pgIdent = regexp.MustCompile(`^[a-z_][a-z0-9_]*$`)
|
|
|
|
func safePGIdent(name string) bool {
|
|
return pgIdent.MatchString(name)
|
|
}
|