From 01852d49f3e63075d2c967ba297e4dab8b905ab2 Mon Sep 17 00:00:00 2001 From: Matt Nolf Date: Tue, 4 Aug 2026 16:26:34 +0100 Subject: [PATCH] refactor(database_observability.postgres): Extract per-database state into dbInstance struct --- .../postgres/component.go | 159 ++++++------------ .../postgres/component_test.go | 50 +++--- .../postgres/instance.go | 91 ++++++++++ 3 files changed, 170 insertions(+), 130 deletions(-) create mode 100644 internal/component/database_observability/postgres/instance.go diff --git a/internal/component/database_observability/postgres/component.go b/internal/component/database_observability/postgres/component.go index 04a2c8c6bbd..1d45fb17860 100644 --- a/internal/component/database_observability/postgres/component.go +++ b/internal/component/database_observability/postgres/component.go @@ -6,7 +6,6 @@ import ( "database/sql" "fmt" "net/http" - "path" "slices" "strings" "sync" @@ -15,10 +14,8 @@ import ( "github.com/lib/pq" pg_collector "github.com/prometheus-community/postgres_exporter/collector" pg_exporter "github.com/prometheus-community/postgres_exporter/exporter" - "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/promhttp" "github.com/prometheus/common/model" - "go.uber.org/atomic" "github.com/grafana/alloy/internal/component" "github.com/grafana/alloy/internal/component/common/loki" @@ -232,20 +229,14 @@ type Collector interface { } type Component struct { - opts component.Options - args Arguments - handler loki.LogsReceiver - fanout *loki.Fanout - mut sync.RWMutex - registry *prometheus.Registry - baseTarget discovery.Target - collectors []Collector - instanceKey string - dbConnection *sql.DB - healthErr *atomic.String - openSQL func(driverName, dataSourceName string) (*sql.DB, error) - logsReceiver loki.LogsReceiver - exporterCollectors []prometheus.Collector + opts component.Options + args Arguments + handler loki.LogsReceiver + fanout *loki.Fanout + mut sync.RWMutex + instance *dbInstance + openSQL func(driverName, dataSourceName string) (*sql.DB, error) + logsReceiver loki.LogsReceiver } func New(opts component.Options, args Arguments) (*Component, error) { @@ -258,23 +249,15 @@ func new(opts component.Options, args Arguments, openFn func(driverName, dataSou args: args, fanout: loki.NewFanout(args.ForwardTo), handler: loki.NewLogsReceiver(), - registry: prometheus.NewRegistry(), - healthErr: atomic.NewString(""), openSQL: openFn, logsReceiver: loki.NewLogsReceiver(), } - instance, err := instanceKey(string(args.DataSourceName)) + instance, err := newDBInstance(opts, string(args.DataSourceName)) if err != nil { return nil, err } - c.instanceKey = instance - - baseTarget, err := c.getBaseTarget() - if err != nil { - return nil, err - } - c.baseTarget = baseTarget + c.instance = instance // Export logs receiver immediately so loki.source.file can wire to it opts.OnStateChange(Exports{ @@ -297,13 +280,13 @@ func (c *Component) Run(ctx context.Context) error { c.mut.Lock() defer c.mut.Unlock() - for _, collector := range c.collectors { + for _, collector := range c.instance.collectors { collector.Stop() } c.cleanupExporterCollectors() - if c.dbConnection != nil { - c.dbConnection.Close() + if c.instance.dbConnection != nil { + c.instance.dbConnection.Close() } }) }() @@ -326,7 +309,7 @@ func (c *Component) Run(ctx context.Context) error { return case <-ticker.C: c.mut.RLock() - hasCollectors := len(c.collectors) > 0 + hasCollectors := len(c.instance.collectors) > 0 c.mut.RUnlock() if !hasCollectors { @@ -344,25 +327,9 @@ func (c *Component) Run(ctx context.Context) error { return nil } -func (c *Component) getBaseTarget() (discovery.Target, error) { - data, err := c.opts.GetServiceData(http_service.ServiceName) - if err != nil { - return discovery.EmptyTarget, fmt.Errorf("failed to get HTTP information: %w", err) - } - httpData := data.(http_service.Data) - - return discovery.NewTargetFromMap(map[string]string{ - model.AddressLabel: httpData.MemoryListenAddr, - model.SchemeLabel: "http", - model.MetricsPathLabel: path.Join(httpData.HTTPPathForComponent(c.opts.ID), "metrics"), - "instance": c.instanceKey, - "job": database_observability.JobName, - }), nil -} - func (c *Component) reportError(errorMsg string, err error) { c.opts.Logger.Error(fmt.Sprintf("%s: %+v", errorMsg, err)) - c.healthErr.Store(fmt.Sprintf("%s: %+v", errorMsg, err)) + c.instance.healthErr.Store(fmt.Sprintf("%s: %+v", errorMsg, err)) } func (c *Component) tryReconnect(ctx context.Context) error { @@ -374,26 +341,26 @@ func (c *Component) tryReconnect(ctx context.Context) error { return err } - c.healthErr.Store("") + c.instance.healthErr.Store("") return nil } // cleanupExporterCollectors releases resources held by embedded exporter collectors. // Callers must hold c.mut. func (c *Component) cleanupExporterCollectors() { - for _, col := range c.exporterCollectors { + for _, col := range c.instance.exporterCollectors { if closable, ok := col.(interface{ CloseServers() }); ok { closable.CloseServers() } - c.registry.Unregister(col) + c.instance.registry.Unregister(col) } - c.exporterCollectors = nil + c.instance.exporterCollectors = nil } func (c *Component) connectAndStartCollectors(ctx context.Context) error { - if c.dbConnection != nil { - c.dbConnection.Close() - c.dbConnection = nil + if c.instance.dbConnection != nil { + c.instance.dbConnection.Close() + c.instance.dbConnection = nil } c.cleanupExporterCollectors() @@ -410,7 +377,7 @@ func (c *Component) connectAndStartCollectors(ctx context.Context) error { dbConnection.Close() return fmt.Errorf("failed to ping database: %w", err) } - c.dbConnection = dbConnection + c.instance.dbConnection = dbConnection rs := dbConnection.QueryRowContext(ctx, selectServerInfo) if err := rs.Err(); err != nil { @@ -471,10 +438,10 @@ func (c *Component) connectAndStartCollectors(ctx context.Context) error { pg_exporter.ExcludeDatabases(c.args.ExcludeDatabases), pg_exporter.WithMetricPrefix("pg"), ) - if err := c.registry.Register(e); err != nil { + if err := c.instance.registry.Register(e); err != nil { return fmt.Errorf("failed to register prometheus_exporter: %w", err) } - c.exporterCollectors = append(c.exporterCollectors, e) + c.instance.exporterCollectors = append(c.instance.exporterCollectors, e) if !exporterArgs.DisableDefaultMetrics { collectorOpts := []pg_collector.Option{pg_collector.WithCollectionTimeout("10s")} @@ -497,14 +464,14 @@ func (c *Component) connectAndStartCollectors(ctx context.Context) error { if err != nil { return fmt.Errorf("failed to create postgres collector: %w", err) } - if err := c.registry.Register(col); err != nil { + if err := c.instance.registry.Register(col); err != nil { return fmt.Errorf("failed to register postgres collector: %w", err) } - c.exporterCollectors = append(c.exporterCollectors, col) + c.instance.exporterCollectors = append(c.instance.exporterCollectors, col) } } - allTargets := append([]discovery.Target{c.baseTarget}, c.args.Targets...) + allTargets := append([]discovery.Target{c.instance.baseTarget}, c.args.Targets...) targets := make([]discovery.Target, 0, len(allTargets)) for _, t := range allTargets { builder := discovery.NewTargetBuilderFrom(t) @@ -518,10 +485,10 @@ func (c *Component) connectAndStartCollectors(ctx context.Context) error { LogsReceiver: c.logsReceiver, }) - for _, collector := range c.collectors { + for _, collector := range c.instance.collectors { collector.Stop() } - c.collectors = nil + c.instance.collectors = nil if err := c.startCollectors(generatedSystemID, engineVersion.String, cp, effectiveExcludeUsers); err != nil { return fmt.Errorf("failed to start collectors: %w", err) @@ -542,7 +509,7 @@ func (c *Component) Update(args component.Arguments) error { return nil } - c.healthErr.Store("") + c.instance.healthErr.Store("") return nil } @@ -579,7 +546,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud startErrors = append(startErrors, errorString) } - entryHandler := addLokiLabels(loki.NewEntryHandler(c.handler.Chan(), func() {}), c.instanceKey, systemID) + entryHandler := addLokiLabels(loki.NewEntryHandler(c.handler.Chan(), func() {}), c.instance.instanceKey, systemID) var tableRegistry *collector.TableRegistry collectors := enableOrDisableCollectors(c.args) @@ -596,7 +563,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud } stCollector, err := collector.NewSchemaDetails(collector.SchemaDetailsArguments{ - DB: c.dbConnection, + DB: c.instance.dbConnection, DSN: string(c.args.DataSourceName), CollectInterval: c.args.SchemaDetailsArguments.CollectInterval, ExcludeDatabases: c.args.ExcludeDatabases, @@ -610,12 +577,12 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := stCollector.Start(context.Background()); err != nil { logStartError(collector.SchemaDetailsCollector, "start", err) } - c.collectors = append(c.collectors, stCollector) + c.instance.collectors = append(c.instance.collectors, stCollector) } if collectors[collector.QueryDetailsCollector] { qCollector, err := collector.NewQueryDetails(collector.QueryDetailsArguments{ - DB: c.dbConnection, + DB: c.instance.dbConnection, CollectInterval: c.args.QueryDetailsArguments.CollectInterval, StatementsLimit: c.args.QueryDetailsArguments.StatementsLimit, ExcludeDatabases: c.args.ExcludeDatabases, @@ -631,7 +598,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := qCollector.Start(context.Background()); err != nil { logStartError(collector.QueryDetailsCollector, "start", err) } - c.collectors = append(c.collectors, qCollector) + c.instance.collectors = append(c.instance.collectors, qCollector) } if collectors[collector.QuerySamplesCollector] { @@ -647,7 +614,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud } aCollector, err := collector.NewQuerySamples(collector.QuerySamplesArguments{ - DB: c.dbConnection, + DB: c.instance.dbConnection, CollectInterval: c.args.QuerySampleArguments.CollectInterval, ExcludeDatabases: c.args.ExcludeDatabases, ExcludeUsers: qsExcludeUsers, @@ -663,16 +630,16 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := aCollector.Start(context.Background()); err != nil { logStartError(collector.QuerySamplesCollector, "start", err) } - c.collectors = append(c.collectors, aCollector) + c.instance.collectors = append(c.instance.collectors, aCollector) } // Connection Info collector is always enabled ciCollector, err := collector.NewConnectionInfo(collector.ConnectionInfoArguments{ DSN: string(c.args.DataSourceName), - Registry: c.registry, + Registry: c.instance.registry, EngineVersion: engineVersion, CloudProvider: cloudProviderInfo, - DB: c.dbConnection, + DB: c.instance.dbConnection, }) if err != nil { logStartError(collector.ConnectionInfoName, "create", err) @@ -681,11 +648,11 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud logStartError(collector.ConnectionInfoName, "start", err) } - c.collectors = append(c.collectors, ciCollector) + c.instance.collectors = append(c.instance.collectors, ciCollector) if collectors[collector.ExplainPlanCollector] { epCollector, err := collector.NewExplainPlan(collector.ExplainPlansArguments{ - DB: c.dbConnection, + DB: c.instance.dbConnection, DSN: string(c.args.DataSourceName), ScrapeInterval: c.args.ExplainPlansArguments.CollectInterval, PerScrapeRatio: c.args.ExplainPlansArguments.PerCollectRatio, @@ -701,12 +668,12 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := epCollector.Start(context.Background()); err != nil { logStartError(collector.ExplainPlanCollector, "start", err) } - c.collectors = append(c.collectors, epCollector) + c.instance.collectors = append(c.instance.collectors, epCollector) } // HealthCheck collector is always enabled hcCollector, err := collector.NewHealthCheck(collector.HealthCheckArguments{ - DB: c.dbConnection, + DB: c.instance.dbConnection, CollectInterval: c.args.HealthCheckArguments.CollectInterval, ExcludeDatabases: c.args.ExcludeDatabases, ExcludeUsers: effectiveExcludeUsers, @@ -719,7 +686,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := hcCollector.Start(context.Background()); err != nil { logStartError(collector.HealthCheckCollector, "start", err) } - c.collectors = append(c.collectors, hcCollector) + c.instance.collectors = append(c.instance.collectors, hcCollector) } // Logs collector is always enabled @@ -727,11 +694,11 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud Receiver: c.logsReceiver, EntryHandler: entryHandler, Logger: c.opts.Logger, - Registry: c.registry, + Registry: c.instance.registry, ExcludeDatabases: c.args.ExcludeDatabases, ExcludeUsers: effectiveExcludeUsers, EnableErrorLogsProcessing: c.args.Logs.EnableErrorLogsProcessing, - DB: c.dbConnection, + DB: c.instance.dbConnection, }) if err != nil { logStartError(collector.LogsCollector, "create", err) @@ -739,7 +706,7 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud if err := logsCollector.Start(context.Background()); err != nil { logStartError(collector.LogsCollector, "start", err) } - c.collectors = append(c.collectors, logsCollector) + c.instance.collectors = append(c.instance.collectors, logsCollector) } if len(startErrors) > 0 { @@ -750,11 +717,11 @@ func (c *Component) startCollectors(systemID string, engineVersion string, cloud } func (c *Component) Handler() http.Handler { - return promhttp.HandlerFor(c.registry, promhttp.HandlerOpts{}) + return promhttp.HandlerFor(c.instance.registry, promhttp.HandlerOpts{}) } func (c *Component) CurrentHealth() component.Health { - if err := c.healthErr.Load(); err != "" { + if err := c.instance.healthErr.Load(); err != "" { return component.Health{ Health: component.HealthTypeUnhealthy, Message: err, @@ -765,7 +732,7 @@ func (c *Component) CurrentHealth() component.Health { var unhealthyCollectors []string c.mut.RLock() - for _, collector := range c.collectors { + for _, collector := range c.instance.collectors { if collector.Stopped() { unhealthyCollectors = append(unhealthyCollectors, collector.Name()) } @@ -787,30 +754,6 @@ func (c *Component) CurrentHealth() component.Health { } } -// instanceKey returns network(hostname:port)/dbname of the Postgres server. -// This is the same key as used by the postgres static integration. -func instanceKey(dsn string) (string, error) { - s, err := collector.ParseURL(dsn) - if err != nil { - return "", fmt.Errorf("cannot parse DSN: %w", err) - } - - // Assign default values to s. - // - // PostgreSQL hostspecs can contain multiple host pairs. We'll assign a host - // and port by default, but otherwise just use the hostname. - if _, ok := s["host"]; !ok { - s["host"] = "localhost" - s["port"] = "5432" - } - - hostport := s["host"] - if p, ok := s["port"]; ok { - hostport += fmt.Sprintf(":%s", p) - } - return fmt.Sprintf("postgresql://%s/%s", hostport, s["dbname"]), nil -} - func addLokiLabels(entryHandler loki.EntryHandler, instanceKey string, systemID string) loki.EntryHandler { entryHandler = loki.AddLabelsMiddleware(model.LabelSet{ "job": database_observability.JobName, diff --git a/internal/component/database_observability/postgres/component_test.go b/internal/component/database_observability/postgres/component_test.go index c8a2e72d568..e736d8b5e7b 100644 --- a/internal/component/database_observability/postgres/component_test.go +++ b/internal/component/database_observability/postgres/component_test.go @@ -52,16 +52,18 @@ func newTestComponent(t *testing.T, openSQL func(string, string) (*sql.DB, error args: args, fanout: loki.NewFanout(args.ForwardTo), handler: loki.NewLogsReceiver(), - registry: prometheus.NewRegistry(), - healthErr: atomic.NewString(""), openSQL: openSQL, logsReceiver: loki.NewLogsReceiver(), + instance: &dbInstance{ + instanceKey: "test-instance", + baseTarget: discovery.NewTargetFromMap(map[string]string{ + "instance": "test-instance", + "job": "database_observability", + }), + registry: prometheus.NewRegistry(), + healthErr: atomic.NewString(""), + }, } - c.instanceKey = "test-instance" - c.baseTarget = discovery.NewTargetFromMap(map[string]string{ - "instance": c.instanceKey, - "job": "database_observability", - }) return c } @@ -917,7 +919,7 @@ func Test_connectAndStartCollectors(t *testing.T) { require.NoError(t, err) // The component should handle nil dbConnection gracefully - assert.Nil(t, c.dbConnection, "dbConnection should be nil initially after failed connection") + assert.Nil(t, c.instance.dbConnection, "dbConnection should be nil initially after failed connection") // Calling connectAndStartCollectors again should not panic err = c.connectAndStartCollectors(context.Background()) @@ -949,14 +951,16 @@ func TestComponent_cleanupExporterCollectors(t *testing.T) { require.NoError(t, registry.Register(collector)) c := &Component{ - registry: registry, - exporterCollectors: []prometheus.Collector{collector}, + instance: &dbInstance{ + registry: registry, + exporterCollectors: []prometheus.Collector{collector}, + }, } c.cleanupExporterCollectors() assert.Equal(t, 1, collector.closeCalls) - assert.Nil(t, c.exporterCollectors) + assert.Nil(t, c.instance.exporterCollectors) assert.False(t, registry.Unregister(collector)) }) @@ -970,13 +974,15 @@ func TestComponent_cleanupExporterCollectors(t *testing.T) { require.NoError(t, registry.Register(collector)) c := &Component{ - registry: registry, - exporterCollectors: []prometheus.Collector{collector}, + instance: &dbInstance{ + registry: registry, + exporterCollectors: []prometheus.Collector{collector}, + }, } c.cleanupExporterCollectors() - assert.Nil(t, c.exporterCollectors) + assert.Nil(t, c.instance.exporterCollectors) assert.False(t, registry.Unregister(collector)) }) } @@ -1001,11 +1007,11 @@ func TestPostgres_Reconnection(t *testing.T) { c, err := New(opts, args) require.NoError(t, err) - c.healthErr.Store("initial error") + c.instance.healthErr.Store("initial error") err = c.tryReconnect(context.Background()) assert.Error(t, err) - assert.NotEmpty(t, c.healthErr.Load()) + assert.NotEmpty(t, c.instance.healthErr.Load()) }) t.Run("tryReconnect succeeds and clears health error", func(t *testing.T) { @@ -1021,7 +1027,7 @@ func TestPostgres_Reconnection(t *testing.T) { // First attempt: connection fails err = c.tryReconnect(context.Background()) assert.Error(t, err) - assert.NotEmpty(t, c.healthErr.Load()) + assert.NotEmpty(t, c.instance.healthErr.Load()) // Second mock: will succeed db2, mock2, err := sqlmock.New(sqlmock.MonitorPingsOption(true), sqlmock.QueryMatcherOption(sqlmock.QueryMatcherEqual)) @@ -1036,14 +1042,14 @@ func TestPostgres_Reconnection(t *testing.T) { // Second attempt: connection succeeds and clears error err = c.tryReconnect(context.Background()) assert.NoError(t, err) - assert.Empty(t, c.healthErr.Load()) + assert.Empty(t, c.instance.healthErr.Load()) }) t.Run("Run exits on context cancellation", func(t *testing.T) { c := newTestComponent(t, func(_, _ string) (*sql.DB, error) { return nil, assert.AnError }) oldCollector := newFakeClosableCollector("test_run_cleanup_old_exporter") - require.NoError(t, c.registry.Register(oldCollector)) - c.exporterCollectors = []prometheus.Collector{oldCollector} + require.NoError(t, c.instance.registry.Register(oldCollector)) + c.instance.exporterCollectors = []prometheus.Collector{oldCollector} ctx, cancel := context.WithCancel(context.Background()) cancel() @@ -1052,8 +1058,8 @@ func TestPostgres_Reconnection(t *testing.T) { assert.NoError(t, err) assert.Equal(t, 1, oldCollector.closeCalls) - assert.Nil(t, c.exporterCollectors) - assert.False(t, c.registry.Unregister(oldCollector)) + assert.Nil(t, c.instance.exporterCollectors) + assert.False(t, c.instance.registry.Unregister(oldCollector)) }) } diff --git a/internal/component/database_observability/postgres/instance.go b/internal/component/database_observability/postgres/instance.go new file mode 100644 index 00000000000..b3cf64b902d --- /dev/null +++ b/internal/component/database_observability/postgres/instance.go @@ -0,0 +1,91 @@ +package postgres + +import ( + "database/sql" + "fmt" + "path" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/common/model" + "go.uber.org/atomic" + + "github.com/grafana/alloy/internal/component" + "github.com/grafana/alloy/internal/component/database_observability" + "github.com/grafana/alloy/internal/component/database_observability/postgres/collector" + "github.com/grafana/alloy/internal/component/discovery" + http_service "github.com/grafana/alloy/internal/service/http" +) + +// dbInstance holds the per-database state of the component: the identity, +// connection, collectors, and metrics registry of a single monitored +// Postgres server. +type dbInstance struct { + instanceKey string + baseTarget discovery.Target + registry *prometheus.Registry + collectors []Collector + dbConnection *sql.DB + healthErr *atomic.String + exporterCollectors []prometheus.Collector +} + +func newDBInstance(opts component.Options, dsn string) (*dbInstance, error) { + key, err := instanceKey(dsn) + if err != nil { + return nil, err + } + + baseTarget, err := getBaseTarget(opts, key) + if err != nil { + return nil, err + } + + return &dbInstance{ + instanceKey: key, + baseTarget: baseTarget, + registry: prometheus.NewRegistry(), + healthErr: atomic.NewString(""), + }, nil +} + +// getBaseTarget returns the scrape target of the component's own metrics +// endpoint for a database instance. +func getBaseTarget(opts component.Options, instanceKey string) (discovery.Target, error) { + data, err := opts.GetServiceData(http_service.ServiceName) + if err != nil { + return discovery.EmptyTarget, fmt.Errorf("failed to get HTTP information: %w", err) + } + httpData := data.(http_service.Data) + + return discovery.NewTargetFromMap(map[string]string{ + model.AddressLabel: httpData.MemoryListenAddr, + model.SchemeLabel: "http", + model.MetricsPathLabel: path.Join(httpData.HTTPPathForComponent(opts.ID), "metrics"), + "instance": instanceKey, + "job": database_observability.JobName, + }), nil +} + +// instanceKey returns network(hostname:port)/dbname of the Postgres server. +// This is the same key as used by the postgres static integration. +func instanceKey(dsn string) (string, error) { + s, err := collector.ParseURL(dsn) + if err != nil { + return "", fmt.Errorf("cannot parse DSN: %w", err) + } + + // Assign default values to s. + // + // PostgreSQL hostspecs can contain multiple host pairs. We'll assign a host + // and port by default, but otherwise just use the hostname. + if _, ok := s["host"]; !ok { + s["host"] = "localhost" + s["port"] = "5432" + } + + hostport := s["host"] + if p, ok := s["port"]; ok { + hostport += fmt.Sprintf(":%s", p) + } + return fmt.Sprintf("postgresql://%s/%s", hostport, s["dbname"]), nil +}