mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-08 06:29:08 +02:00
243 lines
6.6 KiB
Go
243 lines
6.6 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"flag"
|
|
"fmt"
|
|
"log/slog"
|
|
"net/url"
|
|
"os"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
log "github.com/sirupsen/logrus"
|
|
"gorm.io/driver/mysql"
|
|
"gorm.io/driver/postgres"
|
|
"gorm.io/driver/sqlite"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
|
|
"github.com/netbirdio/netbird/idp/dex"
|
|
"github.com/netbirdio/netbird/management/internals/shared/db"
|
|
"github.com/netbirdio/netbird/management/server/activity"
|
|
"github.com/netbirdio/netbird/management/server/store"
|
|
)
|
|
|
|
type config struct {
|
|
mysqlDSN string
|
|
postgresDSN string
|
|
eventsPostgresDSN string
|
|
authPostgresDSN string
|
|
eventsDB string
|
|
authDB string
|
|
mysqlTimezone string
|
|
}
|
|
|
|
type storeMigration struct {
|
|
name string
|
|
source *sql.DB
|
|
dialect dialect
|
|
targetDSN string
|
|
schema func(ctx context.Context, dsn string) error
|
|
skipTables map[string]bool
|
|
}
|
|
|
|
func main() {
|
|
cfg, err := parseFlags(os.Args[1:])
|
|
if errors.Is(err, flag.ErrHelp) {
|
|
return
|
|
}
|
|
if err != nil {
|
|
fmt.Fprintln(os.Stderr, err)
|
|
os.Exit(2)
|
|
}
|
|
|
|
// Silence the store constructors' migration logs.
|
|
log.SetLevel(log.WarnLevel)
|
|
|
|
if err := run(context.Background(), cfg); err != nil {
|
|
fmt.Fprintf(os.Stderr, "migration failed: %v\n", err)
|
|
os.Exit(1)
|
|
}
|
|
}
|
|
|
|
func parseFlags(args []string) (*config, error) {
|
|
cfg := &config{}
|
|
fs := flag.NewFlagSet("netbird-mysql-migrate", flag.ContinueOnError)
|
|
fs.StringVar(&cfg.mysqlDSN, "mysql-dsn", "", "MySQL DSN of the management store (required)")
|
|
fs.StringVar(&cfg.postgresDSN, "postgres-dsn", "", "Postgres DSN to migrate all stores to (required)")
|
|
fs.StringVar(&cfg.eventsPostgresDSN, "events-postgres-dsn", "", "Postgres DSN for activity events (default --postgres-dsn)")
|
|
fs.StringVar(&cfg.authPostgresDSN, "auth-postgres-dsn", "", "Postgres DSN for the embedded IdP (default --postgres-dsn)")
|
|
fs.StringVar(&cfg.eventsDB, "events-db", "/var/lib/netbird/events.db", "SQLite file of the activity events, skipped when missing")
|
|
fs.StringVar(&cfg.authDB, "auth-db", "/var/lib/netbird/idp.db", "SQLite file of the embedded IdP, skipped when missing")
|
|
fs.StringVar(&cfg.mysqlTimezone, "mysql-timezone", "UTC", "time zone the management server ran in, e.g. Europe/Berlin")
|
|
if err := fs.Parse(args); err != nil {
|
|
return nil, err
|
|
}
|
|
if _, err := time.LoadLocation(cfg.mysqlTimezone); err != nil {
|
|
return nil, fmt.Errorf("--mysql-timezone: %w", err)
|
|
}
|
|
|
|
if cfg.mysqlDSN == "" || cfg.postgresDSN == "" {
|
|
return nil, errors.New("--mysql-dsn and --postgres-dsn are required")
|
|
}
|
|
if cfg.eventsPostgresDSN == "" {
|
|
cfg.eventsPostgresDSN = cfg.postgresDSN
|
|
}
|
|
if cfg.authPostgresDSN == "" {
|
|
cfg.authPostgresDSN = cfg.postgresDSN
|
|
}
|
|
return cfg, nil
|
|
}
|
|
|
|
func run(ctx context.Context, cfg *config) error {
|
|
migrations, closeSources, err := openSources(cfg)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer closeSources()
|
|
|
|
for _, m := range migrations {
|
|
if err := m.schema(ctx, m.targetDSN); err != nil {
|
|
return fmt.Errorf("create %s schema: %w", m.name, err)
|
|
}
|
|
}
|
|
|
|
var (
|
|
pools []*pgxpool.Pool
|
|
txs []pgx.Tx
|
|
)
|
|
// Roll back first: pool Close waits for the connections open transactions hold.
|
|
defer func() {
|
|
for _, tx := range txs {
|
|
_ = tx.Rollback(ctx)
|
|
}
|
|
for _, pool := range pools {
|
|
pool.Close()
|
|
}
|
|
}()
|
|
|
|
results := make([]string, 0, len(migrations))
|
|
for _, m := range migrations {
|
|
pool, err := pgxpool.New(ctx, m.targetDSN)
|
|
if err != nil {
|
|
return fmt.Errorf("connect %s target: %w", m.name, err)
|
|
}
|
|
pools = append(pools, pool)
|
|
|
|
tx, err := pool.Begin(ctx)
|
|
if err != nil {
|
|
return fmt.Errorf("begin %s transaction: %w", m.name, err)
|
|
}
|
|
txs = append(txs, tx)
|
|
|
|
tables, rows, err := copyStore(ctx, m.source, m.dialect, tx, m.skipTables)
|
|
if err != nil {
|
|
return fmt.Errorf("%s: %w", m.name, err)
|
|
}
|
|
results = append(results, fmt.Sprintf("%s: %d rows from %d tables", m.name, rows, tables))
|
|
}
|
|
|
|
// Commit last so a failure in any store leaves every target empty.
|
|
for i, tx := range txs {
|
|
if err := tx.Commit(ctx); err != nil {
|
|
return fmt.Errorf("commit %s: %w", migrations[i].name, err)
|
|
}
|
|
}
|
|
|
|
for _, r := range results {
|
|
fmt.Fprintln(os.Stdout, r)
|
|
}
|
|
fmt.Fprintln(os.Stdout, "Done. Point the store, activity and auth store settings at Postgres before starting the server.")
|
|
return nil
|
|
}
|
|
|
|
func openSources(cfg *config) ([]storeMigration, func(), error) {
|
|
var opened []*gorm.DB
|
|
closeAll := func() {
|
|
for _, g := range opened {
|
|
if sqlDB, err := g.DB(); err == nil {
|
|
_ = sqlDB.Close()
|
|
}
|
|
}
|
|
}
|
|
|
|
open := func(dialector gorm.Dialector) (*sql.DB, error) {
|
|
g, err := gorm.Open(dialector, &gorm.Config{Logger: logger.Discard})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
opened = append(opened, g)
|
|
return g.DB()
|
|
}
|
|
|
|
mysqlDB, err := open(mysql.Open(mysqlDSN(cfg.mysqlDSN, cfg.mysqlTimezone)))
|
|
if err != nil {
|
|
closeAll()
|
|
return nil, nil, fmt.Errorf("open MySQL: %w", err)
|
|
}
|
|
migrations := []storeMigration{{
|
|
name: "management", source: mysqlDB, dialect: mysqlDialect,
|
|
targetDSN: cfg.postgresDSN, schema: createManagementSchema,
|
|
}}
|
|
|
|
sqliteStores := []storeMigration{
|
|
{name: "activity", targetDSN: cfg.eventsPostgresDSN, schema: createActivitySchema},
|
|
{name: "auth", targetDSN: cfg.authPostgresDSN, schema: createAuthSchema, skipTables: map[string]bool{"migrations": true}},
|
|
}
|
|
for i, file := range []string{cfg.eventsDB, cfg.authDB} {
|
|
m := sqliteStores[i]
|
|
if _, err := os.Stat(file); errors.Is(err, os.ErrNotExist) {
|
|
fmt.Fprintf(os.Stdout, "%s: %s not found, skipping\n", m.name, file)
|
|
continue
|
|
}
|
|
m.source, err = open(sqlite.Open("file:" + file + "?mode=ro"))
|
|
if err != nil {
|
|
closeAll()
|
|
return nil, nil, fmt.Errorf("open %s: %w", file, err)
|
|
}
|
|
m.dialect = sqliteDialect
|
|
migrations = append(migrations, m)
|
|
}
|
|
|
|
return migrations, closeAll, nil
|
|
}
|
|
|
|
func mysqlDSN(dsn, timezone string) string {
|
|
return strings.TrimSuffix(db.MysqlDSN(dsn), "&loc=Local") + "&loc=" + url.QueryEscape(timezone)
|
|
}
|
|
|
|
func createManagementSchema(ctx context.Context, dsn string) error {
|
|
s, err := store.NewPostgresqlStore(ctx, dsn, nil, false)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return s.Close(ctx)
|
|
}
|
|
|
|
func createActivitySchema(_ context.Context, dsn string) error {
|
|
g, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Discard})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer func() {
|
|
if sqlDB, err := g.DB(); err == nil {
|
|
_ = sqlDB.Close()
|
|
}
|
|
}()
|
|
return g.AutoMigrate(&activity.Event{}, &activity.DeletedUser{})
|
|
}
|
|
|
|
func createAuthSchema(_ context.Context, dsn string) error {
|
|
cfg := dex.Storage{Type: "postgres", Config: map[string]any{"dsn": dsn}}
|
|
s, err := cfg.OpenStorage(slog.New(slog.DiscardHandler))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return s.Close()
|
|
}
|