refactor(migrator): use goose SQL only and auto-import zones on upgrade

移除 Go goose 迁移(bridge/legacy data/zone import),改为 goose SQL 占位
与结构迁移脚本;启动时 Migrate 在删旧列前自动导入历史域名,去掉
migrate-zones 手动命令及文档中的人工导入步骤。
This commit is contained in:
ryan
2026-07-12 15:36:00 +08:00
parent 5b52acdd6c
commit 03b81e5f74
19 changed files with 571 additions and 1252 deletions
@@ -36,8 +36,8 @@ func NormalizeDomains(rawDomains []string) ([]string, error) {
return normalized, nil
}
// DecodeDomains parses legacy route domain fields for the explicit migration
// command. Runtime consumers must read ZoneDomain bindings instead.
// DecodeDomains parses legacy route domain fields for the goose upgrade importer.
// Runtime consumers must read ZoneDomain bindings instead.
func DecodeDomains(raw string, fallbackDomain string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
+298 -118
View File
@@ -5,14 +5,14 @@ package zone
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/routeidentity"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// ImportReport describes the idempotent legacy migration result.
@@ -31,138 +31,244 @@ func (r ImportReport) LogAndReturn(err error) error {
}
type legacyDomain struct {
Domain string
CertID *uint
Remark string
Domain string
CertID *uint
Remark string
ProxyRouteID *uint
}
// legacyRouteRow reads pre-cleanup of_proxy_routes columns via raw scan.
type legacyRouteRow struct {
ID uint `gorm:"column:id"`
Domain string `gorm:"column:domain"`
Domains string `gorm:"column:domains"`
DomainCertIDs string `gorm:"column:domain_cert_ids"`
Remark string `gorm:"column:remark"`
}
// legacyManagedRow reads of_managed_domains while the table still exists.
type legacyManagedRow struct {
Domain string `gorm:"column:domain"`
CertID *uint `gorm:"column:cert_id"`
Remark string `gorm:"column:remark"`
}
// ImportLegacy imports legacy proxy-route / managed-domain rows into Zone tables.
// After the phase-2 schema cleanup, missing legacy columns or tables are skipped.
// ImportLegacyTx imports legacy proxy-route / managed-domain rows into Zone tables
// within an existing SQL transaction (goose runs this on Server upgrade).
// postgres selects $n placeholders; otherwise SQLite-style ? is used.
// Missing legacy columns or tables are skipped so re-runs after phase-2 cleanup are no-ops.
//
//nolint:cyclop // the transactional importer intentionally validates every legacy source in one pass.
func ImportLegacy(ctx context.Context) (report ImportReport, resultErr error) {
conn := db.DB(ctx)
if conn == nil {
return report, fmt.Errorf("database is not initialized")
//nolint:cyclop,gocyclo // single-pass legacy importer validates every source before write.
func ImportLegacyTx(ctx context.Context, tx *sql.Tx, postgres bool) (report ImportReport, err error) {
if tx == nil {
return report, errors.New("transaction is required")
}
resultErr = conn.Transaction(func(tx *gorm.DB) error {
items := make([]legacyDomain, 0)
hasRouteDomains := false
q := func(sqlText string) string { return rebindSQL(sqlText, postgres) }
if tx.Migrator().HasColumn("of_proxy_routes", "domain") &&
tx.Migrator().HasColumn("of_proxy_routes", "domains") {
var routes []legacyRouteRow
if err := tx.Table("of_proxy_routes").
Select("id, domain, domains, domain_cert_ids, remark").
Find(&routes).Error; err != nil {
return err
}
for _, route := range routes {
domains, err := routeidentity.DecodeDomains(route.Domains, route.Domain)
if err != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("route %d: %v", route.ID, err))
continue
}
if len(domains) > 0 {
hasRouteDomains = true
}
certIDs := decodeLegacyCertIDs(route.DomainCertIDs, len(domains))
for i, domain := range domains {
var certID *uint
if i < len(certIDs) && certIDs[i] > 0 {
v := certIDs[i]
certID = &v
}
items = append(items, legacyDomain{Domain: domain, CertID: certID, Remark: route.Remark})
}
items := make([]legacyDomain, 0)
hasRouteDomains := false
hasDomainCol, err := hasTableColumn(ctx, tx, q, postgres, "of_proxy_routes", "domain")
if err != nil {
return report, err
}
hasDomainsCol, err := hasTableColumn(ctx, tx, q, postgres, "of_proxy_routes", "domains")
if err != nil {
return report, err
}
if hasDomainCol && hasDomainsCol {
var collectErr error
items, hasRouteDomains, report.Conflicts, collectErr = collectLegacyRouteDomainsImpl(ctx, tx, q)
if collectErr != nil {
return report, collectErr
}
}
if !hasRouteDomains {
exists, tableErr := hasTable(ctx, tx, q, postgres, "of_managed_domains")
if tableErr != nil {
return report, tableErr
}
if exists {
managed, managedErr := collectLegacyManagedDomains(ctx, tx, q)
if managedErr != nil {
return report, managedErr
}
items = append(items, managed...)
}
}
for _, item := range items {
domain, normErr := normalizeDomain(item.Domain)
if normErr != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", item.Domain, normErr))
continue
}
root, rootErr := zoneRoot(domain)
if rootErr != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", domain, rootErr))
continue
}
if !hasRouteDomains && tx.Migrator().HasTable("of_managed_domains") {
var legacy []legacyManagedRow
if err := tx.Table("of_managed_domains").
Select("domain, cert_id, remark").
Find(&legacy).Error; err != nil {
return err
}
for _, item := range legacy {
items = append(items, legacyDomain(item))
var existingID uint
var existingZoneDomain string
scanErr := tx.QueryRowContext(ctx, q(`
SELECT zd.id, z.domain
FROM of_zone_domains zd
JOIN of_zones z ON z.id = zd.zone_id
WHERE zd.domain = ?
`), domain).Scan(&existingID, &existingZoneDomain)
if scanErr == nil {
if existingZoneDomain != root {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: global domain conflict", domain))
} else if item.ProxyRouteID != nil {
if _, bindErr := tx.ExecContext(ctx, q(`
UPDATE of_zone_domains
SET proxy_route_id = COALESCE(proxy_route_id, ?),
cert_id = COALESCE(cert_id, ?)
WHERE id = ?
`), *item.ProxyRouteID, nullableUint(item.CertID), existingID); bindErr != nil {
return report, bindErr
}
}
continue
}
if !errors.Is(scanErr, sql.ErrNoRows) {
return report, scanErr
}
for _, item := range items {
domain, err := normalizeDomain(item.Domain)
if err != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", item.Domain, err))
continue
}
root, err := zoneRoot(domain)
if err != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", domain, err))
continue
}
var existing model.ZoneDomain
err = tx.Where("domain = ?", domain).First(&existing).Error
if err == nil {
var z model.Zone
if tx.First(&z, existing.ZoneID).Error != nil || z.Domain != root {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: global domain conflict", domain))
}
continue
}
if err != nil && !isNotFound(err) {
return err
}
var zone model.Zone
err = tx.Where("domain = ?", root).First(&zone).Error
if isNotFound(err) {
zone = model.Zone{Domain: root}
if err = tx.Create(&zone).Error; err != nil {
return err
}
report.Zones++
} else if err != nil {
return err
}
if item.CertID != nil {
var cert model.TLSCertificate
if err = tx.First(&cert, *item.CertID).Error; err != nil {
zoneID, zoneErr := ensureZone(ctx, tx, q, root, &report)
if zoneErr != nil {
return report, zoneErr
}
if item.CertID != nil {
var certID uint
if certErr := tx.QueryRowContext(ctx, q(`SELECT id FROM of_tls_certificates WHERE id = ?`), *item.CertID).
Scan(&certID); certErr != nil {
if errors.Is(certErr, sql.ErrNoRows) {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %s", domain, errCertificateNotFound))
continue
}
return report, certErr
}
if err = tx.Create(&model.ZoneDomain{
ZoneID: zone.ID,
Domain: domain,
CertID: item.CertID,
Remark: item.Remark,
}).Error; err != nil {
return err
}
if _, insErr := tx.ExecContext(ctx, q(`
INSERT INTO of_zone_domains (zone_id, proxy_route_id, domain, cert_id, remark, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
`), zoneID, nullableUint(item.ProxyRouteID), domain, nullableUint(item.CertID), item.Remark); insErr != nil {
return report, insErr
}
report.Domains++
}
if len(report.Conflicts) > 0 {
return report, fmt.Errorf("legacy data has conflicts")
}
return report, nil
}
func ensureZone(
ctx context.Context,
tx *sql.Tx,
q func(string) string,
root string,
report *ImportReport,
) (uint, error) {
var zoneID uint
err := tx.QueryRowContext(ctx, q(`SELECT id FROM of_zones WHERE domain = ?`), root).Scan(&zoneID)
if err == nil {
return zoneID, nil
}
if !errors.Is(err, sql.ErrNoRows) {
return 0, err
}
if _, execErr := tx.ExecContext(ctx, q(`
INSERT INTO of_zones (domain, remark, created_at, updated_at)
VALUES (?, '', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
`), root); execErr != nil {
return 0, execErr
}
if err := tx.QueryRowContext(ctx, q(`SELECT id FROM of_zones WHERE domain = ?`), root).Scan(&zoneID); err != nil {
return 0, err
}
report.Zones++
return zoneID, nil
}
func collectLegacyRouteDomainsImpl(
ctx context.Context,
tx *sql.Tx,
q func(string) string,
) (items []legacyDomain, hasRouteDomains bool, conflicts []string, err error) {
// Probe domain_cert_ids: if SELECT fails, fall back without it.
queryWithCert := q(`SELECT id, domain, domains, COALESCE(domain_cert_ids, '[]'), remark FROM of_proxy_routes`)
rows, err := tx.QueryContext(ctx, queryWithCert)
useCert := true
if err != nil {
useCert = false
rows, err = tx.QueryContext(ctx, q(`SELECT id, domain, domains, remark FROM of_proxy_routes`))
if err != nil {
return nil, false, nil, err
}
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var (
id uint
domain string
domains string
certIDs string
remark string
)
if useCert {
if err := rows.Scan(&id, &domain, &domains, &certIDs, &remark); err != nil {
return nil, false, nil, err
}
report.Domains++
} else {
if err := rows.Scan(&id, &domain, &domains, &remark); err != nil {
return nil, false, nil, err
}
certIDs = "[]"
}
if len(report.Conflicts) > 0 {
return fmt.Errorf("legacy data has conflicts")
decoded, decodeErr := routeidentity.DecodeDomains(domains, domain)
if decodeErr != nil {
conflicts = append(conflicts, fmt.Sprintf("route %d: %v", id, decodeErr))
continue
}
return nil
})
return report, resultErr
if len(decoded) > 0 {
hasRouteDomains = true
}
ids := decodeLegacyCertIDs(certIDs, len(decoded))
routeID := id
for i, d := range decoded {
var certID *uint
if i < len(ids) && ids[i] > 0 {
v := ids[i]
certID = &v
}
items = append(items, legacyDomain{
Domain: d,
CertID: certID,
Remark: remark,
ProxyRouteID: &routeID,
})
}
}
return items, hasRouteDomains, conflicts, rows.Err()
}
func collectLegacyManagedDomains(ctx context.Context, tx *sql.Tx, q func(string) string) ([]legacyDomain, error) {
rows, err := tx.QueryContext(ctx, q(`SELECT domain, cert_id, remark FROM of_managed_domains`))
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
items := make([]legacyDomain, 0)
for rows.Next() {
var (
domain string
certID sql.NullInt64
remark string
)
if err := rows.Scan(&domain, &certID, &remark); err != nil {
return nil, err
}
item := legacyDomain{Domain: domain, Remark: remark}
if certID.Valid && certID.Int64 > 0 {
v := uint(certID.Int64)
item.CertID = &v
}
items = append(items, item)
}
return items, rows.Err()
}
func decodeLegacyCertIDs(raw string, count int) []uint {
@@ -176,4 +282,78 @@ func decodeLegacyCertIDs(raw string, count int) []uint {
return values
}
func isNotFound(err error) bool { return err == gorm.ErrRecordNotFound }
func nullableUint(v *uint) any {
if v == nil {
return nil
}
return *v
}
func rebindSQL(query string, postgres bool) string {
if !postgres {
return query
}
var b strings.Builder
b.Grow(len(query) + len(query)/4)
n := 0
for i := 0; i < len(query); i++ {
if query[i] == '?' {
n++
b.WriteByte('$')
b.WriteString(strconv.Itoa(n))
continue
}
b.WriteByte(query[i])
}
return b.String()
}
func hasTable(
ctx context.Context,
tx *sql.Tx,
_ func(string) string,
postgres bool,
table string,
) (bool, error) {
var count int
var err error
if postgres {
err = tx.QueryRowContext(ctx, `
SELECT COUNT(*) FROM information_schema.tables
WHERE table_schema = 'public' AND table_name = $1
`, table).Scan(&count)
} else {
err = tx.QueryRowContext(ctx,
`SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?`, table,
).Scan(&count)
}
if err != nil {
return false, err
}
return count > 0, nil
}
func hasTableColumn(
ctx context.Context,
tx *sql.Tx,
_ func(string) string,
postgres bool,
table, column string,
) (bool, error) {
var count int
var err error
if postgres {
err = tx.QueryRowContext(ctx, `
SELECT COUNT(*) FROM information_schema.columns
WHERE table_schema = 'public' AND table_name = $1 AND column_name = $2
`, table, column).Scan(&count)
} else {
err = tx.QueryRowContext(ctx,
`SELECT COUNT(*) FROM pragma_table_info(?) WHERE name = ?`, table, column,
).Scan(&count)
}
if err != nil {
return false, err
}
return count > 0, nil
}
@@ -0,0 +1,154 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"context"
"database/sql"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupLegacyImportDB(t *testing.T) (*sql.DB, func()) {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
// Pre-phase-2 schema: legacy route columns + managed domains + zone tables.
stmts := []string{
`CREATE TABLE of_zones (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL UNIQUE,
remark TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`,
`CREATE TABLE of_zone_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
zone_id INTEGER NOT NULL,
proxy_route_id INTEGER,
domain TEXT NOT NULL UNIQUE,
cert_id INTEGER,
remark TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`,
`CREATE TABLE of_proxy_routes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
site_name TEXT NOT NULL DEFAULT '',
domain TEXT NOT NULL DEFAULT '',
domains TEXT NOT NULL DEFAULT '[]',
domain_cert_ids TEXT NOT NULL DEFAULT '[]',
origin_url TEXT NOT NULL DEFAULT '',
remark TEXT NOT NULL DEFAULT ''
)`,
`CREATE TABLE of_tls_certificates (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL DEFAULT ''
)`,
`CREATE TABLE of_managed_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL,
cert_id INTEGER,
remark TEXT NOT NULL DEFAULT ''
)`,
}
for _, stmt := range stmts {
_, err := sqlDB.Exec(stmt)
require.NoError(t, err)
}
previous := db.DB(context.Background())
db.SetDB(gormDB)
return sqlDB, func() {
db.SetDB(previous)
_ = sqlDB.Close()
}
}
func TestImportLegacyTxBindsRouteDomains(t *testing.T) {
sqlDB, cleanup := setupLegacyImportDB(t)
defer cleanup()
ctx := context.Background()
_, err := sqlDB.Exec(`INSERT INTO of_tls_certificates (id, name) VALUES (7, 'cert')`)
require.NoError(t, err)
_, err = sqlDB.Exec(`
INSERT INTO of_proxy_routes (id, site_name, domain, domains, domain_cert_ids, origin_url, remark)
VALUES (3, 'api', 'api.example.com', '["api.example.com","www.example.com"]', '[7,7]', 'http://origin', 'r')
`)
require.NoError(t, err)
tx, err := sqlDB.Begin()
require.NoError(t, err)
report, err := ImportLegacyTx(ctx, tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 1, report.Zones)
assert.Equal(t, 2, report.Domains)
var zoneDomain string
require.NoError(t, sqlDB.QueryRow(`SELECT domain FROM of_zones`).Scan(&zoneDomain))
assert.Equal(t, "example.com", zoneDomain)
var count int
require.NoError(t, sqlDB.QueryRow(`SELECT COUNT(*) FROM of_zone_domains WHERE proxy_route_id = 3`).Scan(&count))
assert.Equal(t, 2, count)
// Idempotent re-run
tx, err = sqlDB.Begin()
require.NoError(t, err)
report2, err := ImportLegacyTx(ctx, tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 0, report2.Domains)
}
func TestImportLegacyTxNoOpWithoutLegacyColumns(t *testing.T) {
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
defer sqlDB.Close()
_, err = sqlDB.Exec(`
CREATE TABLE of_zones (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL UNIQUE,
remark TEXT NOT NULL DEFAULT '',
created_at DATETIME, updated_at DATETIME
);
CREATE TABLE of_zone_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
zone_id INTEGER NOT NULL,
proxy_route_id INTEGER,
domain TEXT NOT NULL UNIQUE,
cert_id INTEGER,
remark TEXT NOT NULL DEFAULT '',
created_at DATETIME, updated_at DATETIME
);
CREATE TABLE of_proxy_routes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
site_name TEXT NOT NULL DEFAULT '',
origin_url TEXT NOT NULL DEFAULT ''
);
`)
require.NoError(t, err)
tx, err := sqlDB.Begin()
require.NoError(t, err)
report, err := ImportLegacyTx(context.Background(), tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 0, report.Zones)
assert.Equal(t, 0, report.Domains)
}