File size: 3,420 Bytes
1f10f31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
package common

import (
	"context"
	"database/sql"
	"fmt"
	"log/slog"

	"github.com/XSAM/otelsql"
	"github.com/google/wire"
	"go.opentelemetry.io/otel/metric"
	"go.opentelemetry.io/otel/trace"

	"github.com/openmeterio/openmeter/app/config"
	"github.com/openmeterio/openmeter/openmeter/ent/db"
	"github.com/openmeterio/openmeter/pkg/framework/entutils/entdriver"
	"github.com/openmeterio/openmeter/pkg/framework/pgdriver"
	"github.com/openmeterio/openmeter/tools/migrate"
)

var Database = wire.NewSet(
	wire.Struct(new(Migrator), "*"),

	NewPostgresDriver,
	NewDB,
	NewEntPostgresDriver,
	NewEntClient,
)

// Migrator executes database migrations.
type Migrator struct {
	Config config.PostgresConfig
	Client *db.Client
	Logger *slog.Logger
}

func (m Migrator) Migrate(ctx context.Context) error {
	if !m.Config.AutoMigrate.Enabled() {
		m.Logger.Debug("auto migration is disabled")

		return nil
	}

	m.Logger.Info("running migrations", slog.String("strategy", string(m.Config.AutoMigrate)))

	migrator, err := migrate.New(migrate.MigrateOptions{
		ConnectionString: m.Config.AsURL(),
		Migrations:       migrate.OMMigrationsConfig,
		Logger:           m.Logger,
	})
	if err != nil {
		return fmt.Errorf("failed to create migrator: %w", err)
	}
	defer migrator.CloseOrLogError()

	switch m.Config.AutoMigrate {
	case config.AutoMigrateMigration:
		if err := migrator.Up(); err != nil {
			return fmt.Errorf("failed to migrate db: %w", err)
		}
	case config.AutoMigrateMigrationJob:
		if err := migrator.WaitForMigrationJob(); err != nil {
			return fmt.Errorf("failed to wait for migration job: %w", err)
		}
	}

	m.Logger.Info("database initialized")

	return nil
}

// AdoptLegacyEnt is the explicit upgrade-job entrypoint for databases previously managed by Ent.
func (m Migrator) AdoptLegacyEnt(ctx context.Context) error {
	driver, err := pgdriver.NewPostgresDriver(ctx, m.Config.AsURL())
	if err != nil {
		return fmt.Errorf("open database for legacy Ent adoption: %w", err)
	}
	defer driver.Close()

	return migrate.AdoptLegacyEnt(ctx, driver.DB(), m.Config.AsURL(), m.Logger)
}

func NewPostgresDriver(
	ctx context.Context,
	conf config.PostgresConfig,
	meterProvider metric.MeterProvider,
	meter metric.Meter,
	tracerProvider trace.TracerProvider,
	logger *slog.Logger,
) (*pgdriver.Driver, func(), error) {
	driver, err := pgdriver.NewPostgresDriver(
		ctx,
		conf.AsURL(),
		pgdriver.WithMetricMeter(meter),
		pgdriver.WithTracerProvider(tracerProvider),
		pgdriver.WithMeterProvider(meterProvider),
		pgdriver.WithSpanOptions(otelsql.SpanOptions{
			OmitConnPrepare:      true,
			OmitRows:             true,
			OmitConnectorConnect: true,
		}),
	)
	if err != nil {
		return nil, nil, fmt.Errorf("failed to initialize postgres driver: %w", err)
	}

	return driver, func() {
		err := driver.Close()
		if err != nil {
			logger.Error("failed to close postgres driver", "error", err)
		}
	}, nil
}

// TODO: add closer function?
func NewDB(driver *pgdriver.Driver) *sql.DB {
	return driver.DB()
}

func NewEntPostgresDriver(db *sql.DB, logger *slog.Logger) (*entdriver.EntPostgresDriver, func()) {
	driver := entdriver.NewEntPostgresDriver(db)

	return driver, func() {
		err := driver.Close()
		if err != nil {
			logger.Error("failed to close ent driver", "error", err)
		}
	}
}

// TODO: add closer function?
func NewEntClient(driver *entdriver.EntPostgresDriver) *db.Client {
	return driver.Client()
}