From b8ea423c507519e419c96510edb1902cb47cde4e Mon Sep 17 00:00:00 2001 From: Gavin Frazar Date: Wed, 7 May 2025 16:27:53 -0700 Subject: [PATCH] Add database health checks and routing (#54474) * database service runs health checks on MongoDB, PostgreSQL, and MySQL databases * proxy service routes database user connections based on health status --- api/types/databaseserver.go | 17 +++ api/types/target_health.go | 47 +++++++ api/types/target_health_test.go | 87 +++++++++++++ lib/cloud/mocks/gcp.go | 4 +- lib/srv/db/access_test.go | 7 +- lib/srv/db/common/connect/connect.go | 7 +- lib/srv/db/endpoints/endpoints.go | 98 +++++++++++++++ lib/srv/db/endpoints/endpoints_test.go | 100 +++++++++++++++ lib/srv/db/mongodb/autousers_admin.go | 2 +- lib/srv/db/mongodb/connect.go | 34 +++++- lib/srv/db/mongodb/connect_test.go | 132 ++++++++++++++++++++ lib/srv/db/mysql/engine.go | 70 ++++++++++- lib/srv/db/mysql/engine_test.go | 141 +++++++++++++++++++++ lib/srv/db/postgres/connector.go | 30 +++++ lib/srv/db/postgres/connector_test.go | 77 ++++++++++++ lib/srv/db/server.go | 163 +++++++++++++++++++++++-- lib/srv/db/server_test.go | 116 +++++++++++++++++- lib/srv/heartbeat.go | 7 +- lib/srv/heartbeatv2.go | 12 ++ 19 files changed, 1123 insertions(+), 28 deletions(-) create mode 100644 api/types/target_health_test.go create mode 100644 lib/srv/db/endpoints/endpoints.go create mode 100644 lib/srv/db/endpoints/endpoints_test.go create mode 100644 lib/srv/db/mongodb/connect_test.go create mode 100644 lib/srv/db/postgres/connector_test.go diff --git a/api/types/databaseserver.go b/api/types/databaseserver.go index 47fdd4f224c..ad65057448e 100644 --- a/api/types/databaseserver.go +++ b/api/types/databaseserver.go @@ -57,6 +57,10 @@ type DatabaseServer interface { SetDatabase(Database) error // ProxiedService provides common methods for a proxied service. ProxiedService + // GetTargetHealth returns the database server's target health. + GetTargetHealth() TargetHealth + // SetTargetHealth sets the database server's target health status. + SetTargetHealth(h TargetHealth) } // NewDatabaseServerV3 creates a new database server instance. @@ -285,6 +289,19 @@ func (s *DatabaseServerV3) MatchSearch(values []string) bool { return MatchSearch(nil, values, nil) } +// GetTargetHealth returns the database server's target health. +func (s *DatabaseServerV3) GetTargetHealth() TargetHealth { + if s.Status.TargetHealth == nil { + return TargetHealth{} + } + return *s.Status.TargetHealth +} + +// SetTargetHealth sets the database server's target health status. +func (s *DatabaseServerV3) SetTargetHealth(h TargetHealth) { + s.Status.TargetHealth = &h +} + // DatabaseServers represents a list of database servers. type DatabaseServers []DatabaseServer diff --git a/api/types/target_health.go b/api/types/target_health.go index 193a233b1f2..99731cae763 100644 --- a/api/types/target_health.go +++ b/api/types/target_health.go @@ -16,6 +16,10 @@ limitations under the License. package types +import ( + "time" +) + // TargetHealthProtocol is the network protocol for a health checker. type TargetHealthProtocol string @@ -53,3 +57,46 @@ const ( // encountered an internal error (this is a bug). TargetHealthTransitionReasonInternalError TargetHealthTransitionReason = "internal_error" ) + +// GetTransitionTimestamp returns transition timestamp +func (t *TargetHealth) GetTransitionTimestamp() time.Time { + if t.TransitionTimestamp == nil { + return time.Time{} + } + return *t.TransitionTimestamp +} + +type targetHealthGetter interface { + GetTargetHealth() TargetHealth +} + +// GroupByTargetHealth groups resources by target health and returns [TargetHealthGroups]. +func GroupByTargetHealth[T targetHealthGetter](resources []T) TargetHealthGroups[T] { + var groups TargetHealthGroups[T] + for _, r := range resources { + switch TargetHealthStatus(r.GetTargetHealth().Status) { + case TargetHealthStatusHealthy: + groups.Healthy = append(groups.Healthy, r) + case TargetHealthStatusUnhealthy: + groups.Unhealthy = append(groups.Unhealthy, r) + default: + // all other statuses are equivalent to unknown + groups.Unknown = append(groups.Unknown, r) + } + } + return groups +} + +// TargetHealthGroups holds resources grouped by target health status. +type TargetHealthGroups[T targetHealthGetter] struct { + // Healthy is the resources with [TargetHealthStatusHealthy]. + Healthy []T + // Unhealthy is the resources with [TargetHealthStatusUnhealthy]. + Unhealthy []T + // Unknown is the resources with any status that isn't healthy or unhealthy. + // Namely [TargetHealthStatusUnknown] and the empty string are grouped + // together. + // Agents running with a version prior to health checks will always report + // an empty health status. + Unknown []T +} diff --git a/api/types/target_health_test.go b/api/types/target_health_test.go new file mode 100644 index 00000000000..769c51b095a --- /dev/null +++ b/api/types/target_health_test.go @@ -0,0 +1,87 @@ +/* +Copyright 2025 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package types + +import ( + "fmt" + "math/rand/v2" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGroupByTargetHealth(t *testing.T) { + t.Parallel() + statuses := []TargetHealthStatus{ + TargetHealthStatusHealthy, + TargetHealthStatusUnknown, + TargetHealthStatusUnhealthy, + "", // older agents dont set status + } + + var servers []DatabaseServer + for _, status := range statuses { + for range 10 { + name := fmt.Sprintf("db-%d", len(servers)) + server, err := NewDatabaseServerV3(Metadata{ + Name: name, + }, DatabaseServerSpecV3{ + HostID: "_", + Hostname: "_", + Database: &DatabaseV3{ + Metadata: Metadata{ + Name: name, + }, + Spec: DatabaseSpecV3{ + Protocol: "_", + URI: "_", + AWS: AWS{ + Redshift: Redshift{ + ClusterID: "_", + }, + }, + }, + }, + }) + require.NoError(t, err) + server.SetTargetHealth(TargetHealth{Status: string(status)}) + servers = append(servers, server) + } + } + rand.Shuffle(len(servers), func(i, j int) { + servers[i], servers[j] = servers[j], servers[i] + }) + groups := GroupByTargetHealth(servers) + for _, server := range groups.Healthy { + require.Equal(t, TargetHealthStatusHealthy, + TargetHealthStatus(server.GetTargetHealth().Status), + "server %s is in the wrong group", server.GetName(), + ) + } + for _, server := range groups.Unhealthy { + require.Equal(t, TargetHealthStatusUnhealthy, + TargetHealthStatus(server.GetTargetHealth().Status), + "server %s is in the wrong group", server.GetName(), + ) + } + for _, server := range groups.Unknown { + require.Contains(t, []TargetHealthStatus{TargetHealthStatusUnknown, ""}, + TargetHealthStatus(server.GetTargetHealth().Status), + "server %s is in the wrong group", server.GetName(), + ) + } +} diff --git a/lib/cloud/mocks/gcp.go b/lib/cloud/mocks/gcp.go index 7981a515acc..3d1c1d0e0b8 100644 --- a/lib/cloud/mocks/gcp.go +++ b/lib/cloud/mocks/gcp.go @@ -38,6 +38,8 @@ var _ gcp.SQLAdminClient = (*GCPSQLAdminClientMock)(nil) type GCPSQLAdminClientMock struct { // DatabaseInstance is returned from GetDatabaseInstance. DatabaseInstance *sqladmin.DatabaseInstance + // GetDatabaseInstanceError is returned from GetDatabaseInstance. + GetDatabaseInstanceError error // EphemeralCert is returned from GenerateEphemeralCert. EphemeralCert string // DatabaseUser is returned from GetUser. @@ -56,7 +58,7 @@ func (g *GCPSQLAdminClientMock) UpdateUser(ctx context.Context, db types.Databas } func (g *GCPSQLAdminClientMock) GetDatabaseInstance(ctx context.Context, db types.Database) (*sqladmin.DatabaseInstance, error) { - return g.DatabaseInstance, nil + return g.DatabaseInstance, g.GetDatabaseInstanceError } func (g *GCPSQLAdminClientMock) GenerateEphemeralCert(_ context.Context, _ types.Database, _ time.Time, _ crypto.PublicKey) (string, error) { diff --git a/lib/srv/db/access_test.go b/lib/srv/db/access_test.go index cf22d01335e..0880c07d947 100644 --- a/lib/srv/db/access_test.go +++ b/lib/srv/db/access_test.go @@ -747,7 +747,8 @@ func TestAccessMySQLServerPacket(t *testing.T) { // TestGCPRequireSSL tests connecting to GCP Cloud SQL Postgres and MySQL // databases with an ephemeral client certificate. func TestGCPRequireSSL(t *testing.T) { - ctx := context.Background() + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) user := "alice" testCtx := setupTestContext(ctx, t) testCtx.createUserAndRole(ctx, t, user, "admin", []string{types.Wildcard}, []string{types.Wildcard}) @@ -2267,6 +2268,8 @@ func init() { } func setupTestContext(ctx context.Context, t testing.TB, withDatabases ...withDatabaseOption) *testContext { + ctx, cancel := context.WithCancel(ctx) + t.Cleanup(cancel) testCtx := &testContext{ clusterName: "root.example.com", hostID: uuid.New().String(), @@ -2503,6 +2506,8 @@ func (p *agentParams) setDefaults(c *testContext) { } func (c *testContext) setupDatabaseServer(ctx context.Context, t testing.TB, p agentParams) *Server { + ctx, cancel := context.WithCancel(ctx) + t.Cleanup(cancel) p.setDefaults(c) // Database service credentials. diff --git a/lib/srv/db/common/connect/connect.go b/lib/srv/db/common/connect/connect.go index ee8252f78c3..ebfc8581fdd 100644 --- a/lib/srv/db/common/connect/connect.go +++ b/lib/srv/db/common/connect/connect.go @@ -24,6 +24,7 @@ import ( "log/slog" "math/rand/v2" "net" + "slices" "sort" "strings" @@ -300,10 +301,14 @@ func Connect(ctx context.Context, params ConnectParams) (net.Conn, ConnectStats, return nil, stats, trace.Wrap(err) } + // group the servers by target health, shuffle each group, and then iterate + // over the concatenated groups in order ascending order of health. + params.ShuffleFunc(params.Servers) + groups := types.GroupByTargetHealth(params.Servers) // There may be multiple database servers proxying the same database. If // we get a connection problem error trying to dial one of them, likely // the database server is down so try the next one. - for _, server := range params.ShuffleFunc(params.Servers) { + for _, server := range slices.Concat(groups.Healthy, groups.Unknown, groups.Unhealthy) { stats.attemptedServers++ params.Logger.DebugContext(ctx, "Dialing to database service.", "server", server) tlsConfig, err := GetServerTLSConfig(ctx, ServerTLSConfigParams{ diff --git a/lib/srv/db/endpoints/endpoints.go b/lib/srv/db/endpoints/endpoints.go new file mode 100644 index 00000000000..b03358050eb --- /dev/null +++ b/lib/srv/db/endpoints/endpoints.go @@ -0,0 +1,98 @@ +/* + * Teleport + * Copyright (C) 2025 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package endpoints + +import ( + "context" + "sync" + + "github.com/gravitational/trace" + + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/cloud/gcp" +) + +var ( + resolverBuilders = make(map[string]ResolverBuilder) + resolverBuildersMu sync.RWMutex +) + +// ResolverBuilder constructs a [Resolver]. +type ResolverBuilder func(ctx context.Context, db types.Database, cfg ResolverBuilderConfig) (Resolver, error) + +// Resolver resolves database endpoints. +type Resolver interface { + // Resolve resolves database endpoints. + Resolve(ctx context.Context) ([]string, error) +} + +// ResolverFn is function that implements [Resolver]. +type ResolverFn func(ctx context.Context) ([]string, error) + +// Resolve resolves database endpoints. +func (f ResolverFn) Resolve(ctx context.Context) ([]string, error) { + return f(ctx) +} + +// ResolverBuilderConfig is the config a for a [ResolverBuilder]. +type ResolverBuilderConfig struct { + // GCPClients are clients used to resolve GCP endpoints. + GCPClients GCPClients +} + +// GCPClients are clients used to resolve GCP endpoints. +type GCPClients interface { + // GetGCPSQLAdminClient returns GCP Cloud SQL Admin client. + GetGCPSQLAdminClient(context.Context) (gcp.SQLAdminClient, error) +} + +// RegisterResolver registers a new database endpoint resolver. +func RegisterResolver(builder ResolverBuilder, names ...string) { + resolverBuildersMu.Lock() + defer resolverBuildersMu.Unlock() + for _, name := range names { + resolverBuilders[name] = builder + } +} + +// GetResolver returns a resolver for the given database. +func GetResolver(ctx context.Context, db types.Database, cfg ResolverBuilderConfig) (Resolver, error) { + name := db.GetProtocol() + resolverBuildersMu.RLock() + builder, ok := resolverBuilders[name] + resolverBuildersMu.RUnlock() + if !ok { + return nil, trace.NotFound("database endpoint resolver %q is not registered", name) + } + + resolver, err := builder(ctx, db, cfg) + if err != nil { + return nil, trace.Wrap(err) + } + return resolver, nil +} + +// IsRegistered returns true if the given database protocol has been registered. +func IsRegistered(db types.Database) bool { + name := db.GetProtocol() + resolverBuildersMu.RLock() + defer resolverBuildersMu.RUnlock() + _, ok := resolverBuilders[name] + return ok +} diff --git a/lib/srv/db/endpoints/endpoints_test.go b/lib/srv/db/endpoints/endpoints_test.go new file mode 100644 index 00000000000..fd64d6fe6b6 --- /dev/null +++ b/lib/srv/db/endpoints/endpoints_test.go @@ -0,0 +1,100 @@ +/* + * Teleport + * Copyright (C) 2025 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package endpoints + +import ( + "context" + "fmt" + "testing" + + "github.com/gravitational/trace" + "github.com/stretchr/testify/require" + + "github.com/gravitational/teleport/api/types" +) + +func TestGetResolver(t *testing.T) { + tests := []struct { + desc string + builder ResolverBuilder + wantErr string + }{ + { + desc: "valid", + builder: fakeResolverBuilder{}.builderFunc(), + }, + { + desc: "builder error", + builder: fakeResolverBuilder{builderErr: trace.Errorf("failed to build resolver")}.builderFunc(), + wantErr: "failed to build resolver", + }, + { + desc: "builder not registered", + builder: nil, + wantErr: "is not registered", + }, + } + + ctx := context.Background() + for i, test := range tests { + t.Run(test.desc, func(t *testing.T) { + t.Parallel() + + db := &types.DatabaseV3{} + db.SetName("dummy") + db.Spec.Protocol = fmt.Sprintf("fake-%v", i) + if test.builder != nil { + RegisterResolver(test.builder, db.Spec.Protocol) + t.Cleanup(func() { + resolverBuildersMu.Lock() + defer resolverBuildersMu.Unlock() + delete(resolverBuilders, db.Spec.Protocol) + }) + } + + resolver, err := GetResolver(ctx, db, ResolverBuilderConfig{}) + if test.wantErr != "" { + require.ErrorContains(t, err, test.wantErr) + require.Nil(t, resolver) + return + } + require.NoError(t, err) + require.NotNil(t, resolver) + }) + } +} + +type fakeResolverBuilder struct { + builderErr error +} + +func (f fakeResolverBuilder) builderFunc() ResolverBuilder { + return func(context.Context, types.Database, ResolverBuilderConfig) (Resolver, error) { + if f.builderErr != nil { + return nil, trace.Wrap(f.builderErr) + } + return fakeResolver{}, nil + } +} + +type fakeResolver struct{} + +func (f fakeResolver) Resolve(context.Context) ([]string, error) { + return nil, nil +} diff --git a/lib/srv/db/mongodb/autousers_admin.go b/lib/srv/db/mongodb/autousers_admin.go index cfd242b88d9..d6e61d67d38 100644 --- a/lib/srv/db/mongodb/autousers_admin.go +++ b/lib/srv/db/mongodb/autousers_admin.go @@ -197,7 +197,7 @@ func makeBasicAdminClient(ctx context.Context, sessionCtx *common.Session, e *En return nil, trace.Wrap(err) } - clientCfg, err := makeClientOptionsFromDatabaseURI(sessionCtx) + clientCfg, err := makeClientOptionsFromDatabaseURI(sessionCtx.Database.GetURI()) if err != nil { return nil, trace.Wrap(err) } diff --git a/lib/srv/db/mongodb/connect.go b/lib/srv/db/mongodb/connect.go index 7401b713019..5eda71825f2 100644 --- a/lib/srv/db/mongodb/connect.go +++ b/lib/srv/db/mongodb/connect.go @@ -36,6 +36,7 @@ import ( "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/srv/db/common" + "github.com/gravitational/teleport/lib/srv/db/endpoints" awsutils "github.com/gravitational/teleport/lib/utils/aws" ) @@ -94,7 +95,7 @@ func (e *Engine) connect(ctx context.Context, sessionCtx *common.Session) (drive // getTopologyOptions constructs topology options for connecting to a MongoDB server. func (e *Engine) getTopologyOptions(ctx context.Context, sessionCtx *common.Session) (*topology.Config, description.ServerSelector, error) { - clientCfg, err := makeClientOptionsFromDatabaseURI(sessionCtx) + clientCfg, err := makeClientOptionsFromDatabaseURI(sessionCtx.Database.GetURI()) if err != nil { return nil, nil, trace.Wrap(err) } @@ -216,14 +217,14 @@ func (e *Engine) getAWSAuthenticator(ctx context.Context, sessionCtx *common.Ses return authenticator, nil } -func makeClientOptionsFromDatabaseURI(sessionCtx *common.Session) (*options.ClientOptions, error) { +func makeClientOptionsFromDatabaseURI(uri string) (*options.ClientOptions, error) { clientCfg := options.Client() clientCfg.SetServerSelectionTimeout(common.DefaultMongoDBServerSelectionTimeout) - if strings.HasPrefix(sessionCtx.Database.GetURI(), connstring.SchemeMongoDB) || - strings.HasPrefix(sessionCtx.Database.GetURI(), connstring.SchemeMongoDBSRV) { - clientCfg.ApplyURI(sessionCtx.Database.GetURI()) + if strings.HasPrefix(uri, connstring.SchemeMongoDB) || + strings.HasPrefix(uri, connstring.SchemeMongoDBSRV) { + clientCfg.ApplyURI(uri) } else { - clientCfg.Hosts = []string{sessionCtx.Database.GetURI()} + clientCfg.Hosts = []string{uri} } if err := clientCfg.Validate(); err != nil { return nil, trace.Wrap(err) @@ -231,6 +232,27 @@ func makeClientOptionsFromDatabaseURI(sessionCtx *common.Session) (*options.Clie return clientCfg, nil } +// NewEndpointsResolver returns a health check target endpoint resolver. +// SRV URI (mongodb+srv://) is resolved to a seed list from DNS SRV record +// https://www.mongodb.com/docs/manual/reference/connection-string/#srv-connection-format +func NewEndpointsResolver(_ context.Context, db types.Database, _ endpoints.ResolverBuilderConfig) (endpoints.Resolver, error) { + return newEndpointsResolver(db.GetURI()), nil +} + +func newEndpointsResolver(uri string) endpoints.Resolver { + return endpoints.ResolverFn(func(ctx context.Context) ([]string, error) { + clientCfg, err := makeClientOptionsFromDatabaseURI(uri) + if err != nil { + return nil, trace.Wrap(err) + } + endpoints := make([]string, 0, len(clientCfg.Hosts)) + for _, host := range clientCfg.Hosts { + endpoints = append(endpoints, address.Address(host).String()) + } + return endpoints, nil + }) +} + // getServerSelector returns selector for picking the server to connect to, // which is mostly useful when connecting to a MongoDB replica set. // diff --git a/lib/srv/db/mongodb/connect_test.go b/lib/srv/db/mongodb/connect_test.go new file mode 100644 index 00000000000..21b0c20a73a --- /dev/null +++ b/lib/srv/db/mongodb/connect_test.go @@ -0,0 +1,132 @@ +/* + * Teleport + * Copyright (C) 2025 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package mongodb + +import ( + "context" + "fmt" + "net" + "testing" + + "github.com/gravitational/trace" + "github.com/stretchr/testify/require" + "go.mongodb.org/mongo-driver/x/mongo/driver/dns" +) + +func setDefaultDNSResolver(t *testing.T, resolver *dns.Resolver) { + t.Helper() + // prevent parallel tests from running with the modified DNS resolver. + t.Setenv("withDefaultDNSResolverDisallowsParallelTests", "1") + original := dns.DefaultResolver + t.Cleanup(func() { + dns.DefaultResolver = original + }) + dns.DefaultResolver = resolver +} +func TestNewEndpointsResolver(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + tests := []struct { + desc string + uri string + dnsResolver *dns.Resolver + wantEndpoints []string + wantErr string + }{ + { + desc: "simple host and port", + uri: "mongo1:27017", + wantEndpoints: []string{"mongo1:27017"}, + }, + { + desc: "host without port", + uri: "mongo1", + wantEndpoints: []string{"mongo1:27017"}, + }, + { + desc: "host with scheme", + uri: "mongodb://mongo1:27017", + wantEndpoints: []string{"mongo1:27017"}, + }, + { + desc: "host without scheme", + uri: "mongo1:27017", + wantEndpoints: []string{"mongo1:27017"}, + }, + { + desc: "replicaset", + uri: "mongodb://mongo1:27017,mongo2:27017/?replicaSet=rs0", + wantEndpoints: []string{"mongo1:27017", "mongo2:27017"}, + }, + { + desc: "unresolvable SRV record error", + uri: "mongodb+srv://mongo1.invalid", + wantEndpoints: []string{"mongo1:27017", "mongo2:27017"}, + wantErr: "no such host", + }, + { + desc: "resolvable SRV record", + uri: "mongodb+srv://example.com", + // fake the SRV record lookup for unit test. + // I based the construction of these fake values by doing a real SRV + // record lookup with the real default resolver + dnsResolver: &dns.Resolver{ + LookupSRV: func(svc string, proto string, name string) (string, []*net.SRV, error) { + target := fmt.Sprintf("_%v._%v.%v", svc, proto, name) + records := []*net.SRV{ + { + Target: "foo.com.", + Port: 123, + }, + { + Target: "bar.com.", + Port: 456, + }, + { + Target: "baz.com.", + Port: 789, + }, + } + return target, records, nil + }, + LookupTXT: func(string) ([]string, error) { + // TXT record failures must not break seed list resolution + return nil, trace.Errorf("some error") + }, + }, + wantEndpoints: []string{"foo.com:123", "bar.com:456", "baz.com:789"}, + }, + } + for _, test := range tests { + t.Run(test.desc, func(t *testing.T) { + if test.dnsResolver != nil { + setDefaultDNSResolver(t, test.dnsResolver) + } + resolver := newEndpointsResolver(test.uri) + got, err := resolver.Resolve(ctx) + if test.wantErr != "" { + require.ErrorContains(t, err, test.wantErr) + return + } + require.NoError(t, err) + require.EqualValues(t, test.wantEndpoints, got) + }) + } +} diff --git a/lib/srv/db/mysql/engine.go b/lib/srv/db/mysql/engine.go index 369e4a832ab..0a46b373c82 100644 --- a/lib/srv/db/mysql/engine.go +++ b/lib/srv/db/mysql/engine.go @@ -31,13 +31,16 @@ import ( "github.com/go-mysql-org/go-mysql/packet" "github.com/go-mysql-org/go-mysql/server" "github.com/gravitational/trace" + "golang.org/x/time/rate" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/api/utils/retryutils" + "github.com/gravitational/teleport/lib/cloud/gcp" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/srv/db/cloud" "github.com/gravitational/teleport/lib/srv/db/common" "github.com/gravitational/teleport/lib/srv/db/common/role" + "github.com/gravitational/teleport/lib/srv/db/endpoints" "github.com/gravitational/teleport/lib/srv/db/mysql/protocol" discoverycommon "github.com/gravitational/teleport/lib/srv/discovery/common" "github.com/gravitational/teleport/lib/utils" @@ -495,15 +498,23 @@ func newGCPTLSDialer(tlsConfig *tls.Config) client.Dialer { // by creating a TLS connection to the Cloud Proxy port overriding the // MySQL client's connection. MySQL on the default port does not trust // the ephemeral certificate's CA but Cloud Proxy does. - host, port, err := net.SplitHostPort(address) - if err == nil && port == gcpSQLListenPort { - address = net.JoinHostPort(host, gcpSQLProxyListenPort) - } + address = getGCPTLSAddress(address) tlsDialer := tls.Dialer{Config: tlsConfig} return tlsDialer.DialContext(ctx, network, address) } } +// getGCPTLSAddress returns the appropriate address for a Cloud SQL MySQL +// instance, possibly overriding the default port to instead use the Cloud Proxy +// port. +func getGCPTLSAddress(address string) string { + host, port, err := net.SplitHostPort(address) + if err == nil && port == gcpSQLListenPort { + return net.JoinHostPort(host, gcpSQLProxyListenPort) + } + return address +} + // FetchMySQLVersion connects to MySQL database and tries to read the handshake packet and return the version. // In case of error the message returned by the database is propagated in returned error. func FetchMySQLVersion(ctx context.Context, database types.Database) (string, error) { @@ -525,3 +536,54 @@ const ( // gcpSQLProxyListenPort is the port used by Cloud Proxy for MySQL instances. gcpSQLProxyListenPort = "3307" ) + +// resolverClients are API clients needed to resolve MySQL endpoints. +type resolverClients interface { + // GetGCPSQLAdminClient returns GCP Cloud SQL Admin client. + GetGCPSQLAdminClient(context.Context) (gcp.SQLAdminClient, error) +} + +// NewEndpointsResolver returns a database target endpoint resolver. +func NewEndpointsResolver(_ context.Context, db types.Database, cfg endpoints.ResolverBuilderConfig) (endpoints.Resolver, error) { + return newEndpointsResolver(db, cfg.GCPClients) +} + +func newEndpointsResolver(db types.Database, clients resolverClients) (endpoints.Resolver, error) { + switch { + case db.IsCloudSQL(): + return newCloudSQLEndpointResolver(db, clients), nil + default: + return endpoints.ResolverFn(func(ctx context.Context) ([]string, error) { + return []string{db.GetURI()}, nil + }), nil + } +} + +func newCloudSQLEndpointResolver(db types.Database, clients resolverClients) endpoints.Resolver { + // avoid checking the ssl mode more than once every 15 minutes. + sometimes := rate.Sometimes{Interval: 15 * time.Minute} + var requireSSL bool + return endpoints.ResolverFn(func(ctx context.Context) ([]string, error) { + var requireSSLErr error + sometimes.Do(func() { + clt, err := clients.GetGCPSQLAdminClient(ctx) + if err != nil { + requireSSLErr = trace.Wrap(err) + return + } + + requireSSL, err = cloud.GetGCPRequireSSL(ctx, db, clt) + if err != nil && !trace.IsAccessDenied(err) { + requireSSLErr = trace.Wrap(err) + return + } + }) + if requireSSLErr != nil { + return nil, trace.Wrap(requireSSLErr) + } + if requireSSL { + return []string{getGCPTLSAddress(db.GetURI())}, nil + } + return []string{db.GetURI()}, nil + }) +} diff --git a/lib/srv/db/mysql/engine_test.go b/lib/srv/db/mysql/engine_test.go index f8fba86a61c..b26920d0a57 100644 --- a/lib/srv/db/mysql/engine_test.go +++ b/lib/srv/db/mysql/engine_test.go @@ -29,8 +29,11 @@ import ( "github.com/gravitational/trace" "github.com/stretchr/testify/require" + sqladmin "google.golang.org/api/sqladmin/v1beta4" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/cloud/gcp" + "github.com/gravitational/teleport/lib/cloud/mocks" "github.com/gravitational/teleport/lib/modules" ) @@ -193,3 +196,141 @@ func TestFetchMySQLVersionDoesntBlock(t *testing.T) { require.FailNow(t, "connection should return before") } } + +type fakeDB struct { + types.Database + uri string + isCloudSQL bool +} + +func (f fakeDB) IsCloudSQL() bool { + return f.isCloudSQL +} + +func (f fakeDB) GetGCP() types.GCPCloudSQL { + return types.GCPCloudSQL{ + ProjectID: "", + InstanceID: "", + } +} + +func (f fakeDB) GetURI() string { + return f.uri +} + +type fakeResolverClients struct { + getClientErr error + client *mocks.GCPSQLAdminClientMock +} + +func (f fakeResolverClients) GetGCPSQLAdminClient(context.Context) (gcp.SQLAdminClient, error) { + return f.client, f.getClientErr +} + +func TestNewEndpointsResolver(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + tests := []struct { + desc string + db types.Database + clients resolverClients + wantEndpoints []string + wantErr string + }{ + { + desc: "default", + db: fakeDB{uri: "example.com:3306"}, + wantEndpoints: []string{"example.com:3306"}, + }, + { + desc: "cloudsql db failed to get admin client", + db: fakeDB{uri: "example.com:3306", isCloudSQL: true}, + clients: fakeResolverClients{ + getClientErr: trace.Errorf("failed to get Cloud SQL admin client"), + }, + wantErr: "failed to get Cloud SQL admin client", + }, + { + desc: "cloudsql db requires ssl", + db: fakeDB{uri: "example.com:3306", isCloudSQL: true}, + clients: fakeResolverClients{ + client: &mocks.GCPSQLAdminClientMock{ + DatabaseInstance: &sqladmin.DatabaseInstance{ + Settings: &sqladmin.Settings{ + IpConfiguration: &sqladmin.IpConfiguration{ + RequireSsl: true, + }, + }, + }, + }, + }, + wantEndpoints: []string{"example.com:3307"}, + }, + { + desc: "cloudsql db non-default port is not overridden when ssl is required", + db: fakeDB{uri: "example.com:1234", isCloudSQL: true}, + clients: fakeResolverClients{ + client: &mocks.GCPSQLAdminClientMock{ + DatabaseInstance: &sqladmin.DatabaseInstance{ + Settings: &sqladmin.Settings{ + IpConfiguration: &sqladmin.IpConfiguration{ + RequireSsl: true, + }, + }, + }, + }, + }, + wantEndpoints: []string{"example.com:1234"}, + }, + { + desc: "cloudsql db does not require ssl", + db: fakeDB{uri: "example.com:3306", isCloudSQL: true}, + clients: fakeResolverClients{ + client: &mocks.GCPSQLAdminClientMock{ + DatabaseInstance: &sqladmin.DatabaseInstance{ + Settings: &sqladmin.Settings{ + IpConfiguration: &sqladmin.IpConfiguration{ + RequireSsl: false, + }, + }, + }, + }, + }, + wantEndpoints: []string{"example.com:3306"}, + }, + { + desc: "cloudsql db access denied to check ssl setting ", + db: fakeDB{uri: "example.com:3306", isCloudSQL: true}, + clients: fakeResolverClients{ + client: &mocks.GCPSQLAdminClientMock{ + GetDatabaseInstanceError: trace.AccessDenied("unauthorized"), + }, + }, + wantEndpoints: []string{"example.com:3306"}, + }, + { + desc: "cloudsql db other error when checking ssl setting ", + db: fakeDB{uri: "example.com:3306", isCloudSQL: true}, + clients: fakeResolverClients{ + client: &mocks.GCPSQLAdminClientMock{ + GetDatabaseInstanceError: trace.NotFound("not found"), + }, + }, + wantErr: "Failed to get Cloud SQL instance information", + }, + } + for _, test := range tests { + t.Run(test.desc, func(t *testing.T) { + resolver, err := newEndpointsResolver(test.db, test.clients) + require.NoError(t, err) + got, err := resolver.Resolve(ctx) + if test.wantErr != "" { + require.ErrorContains(t, err, test.wantErr) + return + } + require.NoError(t, err) + require.EqualValues(t, test.wantEndpoints, got) + }) + } +} diff --git a/lib/srv/db/postgres/connector.go b/lib/srv/db/postgres/connector.go index 3a73d9e58b0..dae31961ab6 100644 --- a/lib/srv/db/postgres/connector.go +++ b/lib/srv/db/postgres/connector.go @@ -20,6 +20,9 @@ import ( "context" "fmt" "log/slog" + "net" + "slices" + "strconv" "time" "github.com/gravitational/trace" @@ -30,6 +33,7 @@ import ( libcloud "github.com/gravitational/teleport/lib/cloud" "github.com/gravitational/teleport/lib/srv/db/cloud" "github.com/gravitational/teleport/lib/srv/db/common" + "github.com/gravitational/teleport/lib/srv/db/endpoints" discoverycommon "github.com/gravitational/teleport/lib/srv/discovery/common" ) @@ -45,6 +49,32 @@ type connector struct { startupParams map[string]string } +// NewEndpointsResolver returns a health check target endpoint resolver. +func NewEndpointsResolver(_ context.Context, db types.Database, _ endpoints.ResolverBuilderConfig) (endpoints.Resolver, error) { + return newEndpointsResolver(db.GetURI()) +} + +func newEndpointsResolver(uri string) (endpoints.Resolver, error) { + config, err := pgconn.ParseConfig(fmt.Sprintf("postgres://%s", uri)) + if err != nil { + return nil, trace.Wrap(err) + } + addrs := make([]string, 0, len(config.Fallbacks)+1) + hostPort := net.JoinHostPort(config.Host, strconv.Itoa(int(config.Port))) + addrs = append(addrs, hostPort) + for _, fb := range config.Fallbacks { + hostPort := net.JoinHostPort(fb.Host, strconv.Itoa(int(fb.Port))) + // pgconn duplicates the host/port in its fallbacks for some reason, so + // we de-duplicate and preserve the fallback order + if !slices.Contains(addrs, hostPort) { + addrs = append(addrs, hostPort) + } + } + return endpoints.ResolverFn(func(context.Context) ([]string, error) { + return addrs, nil + }), nil +} + func (c *connector) getConnectConfig(ctx context.Context) (*pgconn.Config, error) { // The driver requires the config to be built by parsing the connection // string so parse the basic template and then fill in the rest of diff --git a/lib/srv/db/postgres/connector_test.go b/lib/srv/db/postgres/connector_test.go new file mode 100644 index 00000000000..66a3458ba5b --- /dev/null +++ b/lib/srv/db/postgres/connector_test.go @@ -0,0 +1,77 @@ +/* + * Teleport + * Copyright (C) 2025 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package postgres + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNewEndpointsResolver(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + tests := []struct { + desc string + uri string + wantEndpoints []string + wantErr string + }{ + { + // we don't allow invalid URIs anyway + desc: "URI parse fail", + uri: "host1,host2:123,host3", + wantErr: "failed to parse", + }, + { + desc: "single endpoint", + uri: "example.com:5432", + wantEndpoints: []string{"example.com:5432"}, + }, + { + desc: "single endpoint custom port", + uri: "example.com:123", + wantEndpoints: []string{"example.com:123"}, + }, + { + desc: "multiple endpoints", + uri: "host1,host2/somedb?target_session_attrs=any&application_name=myapp", + wantEndpoints: []string{"host1:5432", "host2:5432"}, + }, + { + desc: "multiple endpoints custom ports", + uri: "host1,host2:456/somedb?target_session_attrs=any&application_name=myapp", + wantEndpoints: []string{"host1:456", "host2:456"}, + }, + } + for _, test := range tests { + t.Run(test.desc, func(t *testing.T) { + resolver, err := newEndpointsResolver(test.uri) + if test.wantErr != "" { + require.ErrorContains(t, err, test.wantErr) + return + } + got, err := resolver.Resolve(ctx) + require.NoError(t, err) + require.EqualValues(t, test.wantEndpoints, got) + }) + } +} diff --git a/lib/srv/db/server.go b/lib/srv/db/server.go index 9cd0703ef16..4ac0cd4ffa8 100644 --- a/lib/srv/db/server.go +++ b/lib/srv/db/server.go @@ -44,6 +44,7 @@ import ( "github.com/gravitational/teleport/lib/cloud/awsconfig" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/events" + "github.com/gravitational/teleport/lib/healthcheck" "github.com/gravitational/teleport/lib/inventory" "github.com/gravitational/teleport/lib/inventory/metadata" "github.com/gravitational/teleport/lib/labels" @@ -59,6 +60,7 @@ import ( "github.com/gravitational/teleport/lib/srv/db/common" "github.com/gravitational/teleport/lib/srv/db/dynamodb" "github.com/gravitational/teleport/lib/srv/db/elasticsearch" + "github.com/gravitational/teleport/lib/srv/db/endpoints" "github.com/gravitational/teleport/lib/srv/db/mongodb" "github.com/gravitational/teleport/lib/srv/db/mysql" "github.com/gravitational/teleport/lib/srv/db/objects" @@ -71,6 +73,7 @@ import ( discoverycommon "github.com/gravitational/teleport/lib/srv/discovery/common" "github.com/gravitational/teleport/lib/srv/discovery/fetchers/db" "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/lib/utils/log" ) func init() { @@ -89,6 +92,10 @@ func init() { common.RegisterEngine(spanner.NewEngine, defaults.ProtocolSpanner) objects.RegisterObjectFetcher(postgres.NewObjectFetcher, defaults.ProtocolPostgres) + + endpoints.RegisterResolver(postgres.NewEndpointsResolver, defaults.ProtocolPostgres, defaults.ProtocolCockroachDB) + endpoints.RegisterResolver(mysql.NewEndpointsResolver, defaults.ProtocolMySQL) + endpoints.RegisterResolver(mongodb.NewEndpointsResolver, defaults.ProtocolMongoDB) } // Config is the configuration for a database proxy server. @@ -168,6 +175,8 @@ type Config struct { // getEngineFn returns a [common.Engine]. It can be overridden in tests to // customize the returned engine. getEngineFn func(types.Database, common.EngineConfig) (common.Engine, error) + // healthCheckManager manages registered health checks for databases. + healthCheckManager healthcheck.Manager } // NewAuditFn defines a function that creates an audit logger. @@ -326,6 +335,18 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) { c.ShutdownPollPeriod = defaults.ShutdownPollPeriod } + if c.healthCheckManager == nil { + manager, err := healthcheck.NewManager(ctx, healthcheck.ManagerConfig{ + Component: teleport.ComponentDatabase, + Events: c.AccessPoint, + HealthCheckConfigReader: c.AccessPoint, + }) + if err != nil { + return trace.Wrap(err, "failed to start database health check manager") + } + c.healthCheckManager = manager + } + return nil } @@ -486,6 +507,12 @@ func New(ctx context.Context, config Config) (*Server, error) { // startDatabase performs initialization actions for the provided database // such as starting dynamic labels and initializing CA certificate. func (s *Server) startDatabase(ctx context.Context, database types.Database) error { + if err := s.startHealthCheck(ctx, database); err != nil { + s.log.DebugContext(ctx, "Failed to start database health checker", + "db", database.GetName(), + "error", err, + ) + } // For cloud-hosted databases (RDS, Redshift, GCP), try to automatically // download a CA certificate. // TODO(r0mant): This should ideally become a part of cloud metadata service. @@ -541,16 +568,39 @@ func (s *Server) startDatabase(ctx context.Context, database types.Database) err } // stopDatabase uninitializes the database with the specified name. -func (s *Server) stopDatabase(ctx context.Context, name string) error { +func (s *Server) stopDatabase(ctx context.Context, db types.Database) error { // Stop database object importer. - if err := s.cfg.DatabaseObjects.StopImporter(name); err != nil { - s.log.WarnContext(ctx, "Failed to stop database object importer.", "db", name, "error", err) + if err := s.cfg.DatabaseObjects.StopImporter(db.GetName()); err != nil { + s.log.WarnContext(ctx, "Failed to stop database object importer", + "db", log.StringerAttr(db), + "error", err, + ) } - s.stopDynamicLabels(name) - if err := s.stopHeartbeat(name); err != nil { - return trace.Wrap(err) + s.stopDynamicLabels(db.GetName()) + + var errors []error + if err := s.stopHeartbeat(db.GetName()); err != nil { + s.log.WarnContext(ctx, "Failed to stop database heartbeat", + "db", log.StringerAttr(db), + "error", err, + ) + errors = append(errors, err) } - s.log.DebugContext(ctx, "Stopped database.", "db", name) + + if err := s.stopHealthCheck(db); err != nil { + s.log.WarnContext(ctx, "Failed to stop database health checker", + "db", log.StringerAttr(db), + "error", err, + ) + errors = append(errors, err) + } + + if len(errors) > 0 { + return trace.NewAggregate(errors...) + } + s.log.DebugContext(ctx, "Stopped database", + "db", log.StringerAttr(db), + ) return nil } @@ -611,7 +661,7 @@ func (s *Server) stopDynamicLabels(name string) { func (s *Server) registerDatabase(ctx context.Context, database types.Database) error { if err := s.startDatabase(ctx, database); err != nil { // Cleanup in case database was initialized only partially. - if errStop := s.stopDatabase(ctx, database.GetName()); errStop != nil { + if errStop := s.stopDatabase(ctx, database); errStop != nil { return trace.NewAggregate(err, errStop) } return trace.Wrap(err) @@ -625,7 +675,7 @@ func (s *Server) registerDatabase(ctx context.Context, database types.Database) // updateDatabase updates database that is already registered. func (s *Server) updateDatabase(ctx context.Context, database types.Database) error { // Stop heartbeat and dynamic labels before starting new ones. - if err := s.stopDatabase(ctx, database.GetName()); err != nil { + if err := s.stopDatabase(ctx, database); err != nil { return trace.Wrap(err) } if err := s.registerDatabase(ctx, database); err != nil { @@ -656,7 +706,7 @@ func (s *Server) unregisterDatabase(ctx context.Context, database types.Database // unregisters it from the list of proxied databases. func (s *Server) stopProxyingAndDeleteDatabase(ctx context.Context, database types.Database) error { // Stop heartbeat and dynamic labels updates. - if err := s.stopDatabase(ctx, database.GetName()); err != nil { + if err := s.stopDatabase(ctx, database); err != nil { return trace.Wrap(err) } // Heartbeat is stopped but if we don't remove this database server, @@ -782,6 +832,7 @@ func (s *Server) getServerInfo(ctx context.Context, database types.Database) (*t Database: copy, ProxyIDs: s.cfg.ConnectedProxyGetter.GetProxyIDs(), }) + server.SetTargetHealth(s.getTargetHealth(ctx, database)) return server, trace.Wrap(err) } @@ -799,6 +850,12 @@ func (s *Server) getRotationState() types.Rotation { // Start starts proxying all server's registered databases. func (s *Server) Start(ctx context.Context) (err error) { + // Start the health check manager that will be monitoring database + // connection health. + if err := s.cfg.healthCheckManager.Start(ctx); err != nil { + return trace.Wrap(err) + } + // Start IAM service that will be configuring IAM auth for databases. if err := s.cfg.CloudIAM.Start(ctx); err != nil { return trace.Wrap(err) @@ -996,6 +1053,14 @@ func (s *Server) close(ctx context.Context) error { s.log.WarnContext(ctx, "Deleting all databases failed", "error", err) } + if s.cfg.healthCheckManager != nil { + if err := s.cfg.healthCheckManager.Close(); err != nil { + s.log.WarnContext(ctx, "Closing health check manager failed", + "error", err, + ) + } + } + s.mu.Lock() clear(s.proxiedDatabases) clear(s.dynamicLabels) @@ -1352,3 +1417,81 @@ func (s *Server) trackSession(ctx context.Context, sessionCtx *common.Session) e return nil } + +// startHealthCheck starts health checks for the database. +func (s *Server) startHealthCheck(ctx context.Context, db types.Database) error { + resolver, err := s.getEndpointsResolver(ctx, db) + if err != nil { + return trace.Wrap(err) + } + err = s.cfg.healthCheckManager.AddTarget(healthcheck.Target{ + GetResource: func() types.ResourceWithLabels { + s.mu.RLock() + defer s.mu.RUnlock() + return s.copyDatabaseWithUpdatedLabelsLocked(db) + }, + ResolverFn: resolver, + }) + return trace.Wrap(err) +} + +// stopHealthCheck stops health checks for the database. +func (s *Server) stopHealthCheck(db types.Database) error { + err := s.cfg.healthCheckManager.RemoveTarget(db) + if err != nil && !trace.IsNotFound(err) { + // not found shouldn't happen, but we can ignore it in any case + return trace.Wrap(err) + } + return nil +} + +// getTargetHealth returns the target health for the database. +func (s *Server) getTargetHealth(ctx context.Context, db types.Database) types.TargetHealth { + if err := checkSupportsHealthChecks(db); err != nil { + return types.TargetHealth{} + } + + health, err := s.cfg.healthCheckManager.GetTargetHealth(db) + if err == nil { + return *health + } + if trace.IsNotFound(err) { + return types.TargetHealth{ + Status: string(types.TargetHealthStatusUnknown), + TransitionReason: string(types.TargetHealthTransitionReasonDisabled), + Message: "The target health checker was not found", + } + } + + s.log.WarnContext(ctx, "Failed to get database target endpoint health", + "db", db.String(), + "error", err, + ) + return types.TargetHealth{ + Status: string(types.TargetHealthStatusUnknown), + TransitionReason: string(types.TargetHealthTransitionReasonInternalError), + TransitionError: err.Error(), + Message: "The database service failed to get the database target endpoint health status (this is a bug)", + } +} + +// getEndpointsResolver gets a health check endpoint resolver for the database. +func (s *Server) getEndpointsResolver(ctx context.Context, db types.Database) (healthcheck.EndpointsResolverFunc, error) { + resolver, err := endpoints.GetResolver(ctx, db, endpoints.ResolverBuilderConfig{ + GCPClients: s.cfg.CloudClients, + }) + if err != nil { + return nil, trace.Wrap(err) + } + return resolver.Resolve, nil +} + +// checkSupportsHealthChecks returns nil if the database supports health checks. +// TODO(gavin): add resolvers for all database protocols that we support, then +// remove this helper. +func checkSupportsHealthChecks(db types.Database) error { + if !endpoints.IsRegistered(db) { + return trace.NotImplemented("endpoint health checks for database protocol %q are not supported", db.GetProtocol()) + } + return nil +} diff --git a/lib/srv/db/server_test.go b/lib/srv/db/server_test.go index 5f16e3b46d2..4f8bd34fb74 100644 --- a/lib/srv/db/server_test.go +++ b/lib/srv/db/server_test.go @@ -20,6 +20,7 @@ package db import ( "context" + "encoding/pem" "fmt" "io" "sync/atomic" @@ -33,14 +34,22 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.mongodb.org/mongo-driver/mongo" + sqladmin "google.golang.org/api/sqladmin/v1beta4" + "google.golang.org/protobuf/types/known/durationpb" "github.com/gravitational/teleport/api/client/proto" apidefaults "github.com/gravitational/teleport/api/defaults" + healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1" + labelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/label/v1" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/types/healthcheckconfig" "github.com/gravitational/teleport/api/utils/retryutils" + "github.com/gravitational/teleport/lib/cloud/mocks" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/limiter" + "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/srv/db/common" "github.com/gravitational/teleport/lib/srv/db/mysql" ) @@ -313,8 +322,9 @@ func TestHeartbeatEvents(t *testing.T) { t.Run(name, func(t *testing.T) { var heartbeatEvents int64 heartbeatRecorder := func(err error) { - require.NoError(t, err) + assert.NoError(t, err) atomic.AddInt64(&heartbeatEvents, 1) + assert.LessOrEqual(t, atomic.LoadInt64(&heartbeatEvents), expectedHeartbeatCount(test.staticDatabases)) } testCtx := setupTestContext(ctx, t) @@ -636,3 +646,107 @@ func databaseServerWithActiveConnection(t *testing.T, ctx context.Context) (*Ser return testCtx.server, connErrCh, cancelConn } + +func TestHealthCheck(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + user := "alice" + testCtx := setupTestContext(ctx, t) + testCtx.createUserAndRole(ctx, t, user, "admin", []string{types.Wildcard}, []string{types.Wildcard}) + + hcc := newHealthCheckConfig(t, "match-all") + _, err := testCtx.authServer.CreateHealthCheckConfig(ctx, hcc) + require.NoError(t, err) + require.EventuallyWithT(t, func(t *assert.CollectT) { + _, err := testCtx.authServer.GetHealthCheckConfig(ctx, "match-all") + assert.NoError(t, err) + }, time.Second, time.Millisecond*100, "waiting for health check config") + + // Generate ephemeral cert returned from mock GCP API. + ephemeralCert, err := common.MakeTestClientTLSCert(common.TestClientConfig{ + AuthClient: testCtx.authClient, + AuthServer: testCtx.authServer, + Cluster: testCtx.clusterName, + Username: user, + }) + require.NoError(t, err) + certPEM := pem.EncodeToMemory(&pem.Block{ + Type: "CERTIFICATE", + Bytes: ephemeralCert.Certificate[0], + }) + + testCtx.server = testCtx.setupDatabaseServer(ctx, t, agentParams{ + Databases: []types.Database{ + withCloudSQLMySQLTLS("cloudsql-mysql", user, cloudSQLPassword)(t, ctx, testCtx), + withCloudSQLPostgres("cloudsql-postgres", cloudSQLAuthToken)(t, ctx, testCtx), + withSelfHostedMongo("self-hosted-mongo")(t, ctx, testCtx), + withSelfHostedMySQL("self-hosted-mysql")(t, ctx, testCtx), + withSelfHostedPostgres("self-hosted-postgres")(t, ctx, testCtx), + }, + GCPSQL: &mocks.GCPSQLAdminClientMock{ + EphemeralCert: string(certPEM), + DatabaseInstance: &sqladmin.DatabaseInstance{ + Settings: &sqladmin.Settings{ + IpConfiguration: &sqladmin.IpConfiguration{ + RequireSsl: true, + }, + }, + }, + }, + }) + go testCtx.startHandlingConnections() + for _, db := range testCtx.server.cfg.Databases { + t.Run(db.GetName(), func(t *testing.T) { + t.Parallel() + waitForHealthStatus(t, ctx, db.GetName(), testCtx.authServer, types.TargetHealthStatusHealthy) + }) + } +} + +func waitForHealthStatus(t *testing.T, ctx context.Context, name string, serverGetter services.DatabaseServersGetter, want types.TargetHealthStatus) { + t.Helper() + timeout := 15 * time.Second + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + require.EventuallyWithT(t, func(t *assert.CollectT) { + servers, err := serverGetter.GetDatabaseServers(ctx, apidefaults.Namespace) + if !assert.NoError(t, err) { + return + } + var server types.DatabaseServer + for _, s := range servers { + if s.GetName() == name { + server = s + break + } + } + if server == nil { + assert.FailNowf(t, "failed to find db_server", "db_server=%s", name) + return + } + health := server.GetTargetHealth() + assert.Equal(t, string(want), health.Status) + }, timeout, time.Millisecond*250, "waiting for database %s to become healthy", name) +} + +func newHealthCheckConfig(t *testing.T, name string) *healthcheckconfigv1.HealthCheckConfig { + t.Helper() + out, err := healthcheckconfig.NewHealthCheckConfig(name, + &healthcheckconfigv1.HealthCheckConfigSpec{ + Match: &healthcheckconfigv1.Matcher{ + DbLabels: []*labelv1.Label{{ + Name: types.Wildcard, + Values: []string{types.Wildcard}, + }}, + }, + Interval: durationpb.New(apidefaults.HealthCheckInterval), + Timeout: durationpb.New(apidefaults.HealthCheckTimeout), + HealthyThreshold: 1, + UnhealthyThreshold: 1, + }, + ) + require.NoError(t, err) + return out +} diff --git a/lib/srv/heartbeat.go b/lib/srv/heartbeat.go index a9b9ab19a10..648fff06372 100644 --- a/lib/srv/heartbeat.go +++ b/lib/srv/heartbeat.go @@ -31,6 +31,7 @@ import ( "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/auth/authclient" "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/utils/log" ) // HeartbeatI abstracts over the basic interface of Heartbeat and HeartbeatV2. This can be removed @@ -162,9 +163,9 @@ func NewHeartbeat(cfg HeartbeatConfig) (*Heartbeat, error) { } h.logger.DebugContext(ctx, "Starting heartbeat with announce period", "mode", cfg.Mode, - "keep_alive_period", cfg.KeepAlivePeriod, - "announce_period", cfg.AnnouncePeriod, - "check_period", cfg.CheckPeriod, + "keep_alive_period", log.StringerAttr(cfg.KeepAlivePeriod), + "announce_period", log.StringerAttr(cfg.AnnouncePeriod), + "check_period", log.StringerAttr(cfg.CheckPeriod), ) return h, nil } diff --git a/lib/srv/heartbeatv2.go b/lib/srv/heartbeatv2.go index 299c1b274e0..460d1e10ad6 100644 --- a/lib/srv/heartbeatv2.go +++ b/lib/srv/heartbeatv2.go @@ -401,6 +401,11 @@ func (h *HeartbeatV2) run() { case <-h.closeContext.Done(): return } + + // check if we are closing to avoid randomly looping back into the sender + if h.closing() { + return + } } } @@ -502,9 +507,16 @@ func (h *HeartbeatV2) onHeartbeat(err error) { if h.onHeartbeatInner == nil { return } + if h.closing() { + return + } h.onHeartbeatInner(err) } +func (h *HeartbeatV2) closing() bool { + return h.closeContext.Err() != nil +} + // heartbeatV2Driver is the pluggable core of the HeartbeatV2 type. A service needing to use HeartbeatV2 should // have a corresponding implementation of heartbeatV2Driver. type heartbeatV2Driver interface {