From 559ea705f5981ad8beb2cab7e1699592d4f33e41 Mon Sep 17 00:00:00 2001 From: Nic Klaassen Date: Tue, 17 Sep 2024 13:07:36 -0700 Subject: [PATCH] eliminate hardcoded RSA from lib/srv/db (#46660) --- lib/cloud/gcp/sql.go | 57 +++++++++++++++----------------- lib/cloud/mocks/gcp.go | 6 ++-- lib/srv/db/access_test.go | 7 +++- lib/srv/db/auth_test.go | 7 ++++ lib/srv/db/cloud/gcp.go | 27 ++++++++++++--- lib/srv/db/common/auth.go | 23 +++++++++++-- lib/srv/db/common/test.go | 3 +- lib/srv/db/mysql/engine.go | 8 ++++- lib/srv/db/postgres/connector.go | 8 ++++- 9 files changed, 102 insertions(+), 44 deletions(-) diff --git a/lib/cloud/gcp/sql.go b/lib/cloud/gcp/sql.go index 109d8628891..af6ceffe798 100644 --- a/lib/cloud/gcp/sql.go +++ b/lib/cloud/gcp/sql.go @@ -20,9 +20,8 @@ package gcp import ( "context" - "crypto/rand" + "crypto" "crypto/rsa" - "crypto/tls" "crypto/x509" "encoding/pem" "fmt" @@ -31,7 +30,6 @@ import ( "github.com/gravitational/trace" sqladmin "google.golang.org/api/sqladmin/v1beta4" - "github.com/gravitational/teleport/api/constants" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/api/utils/keys" ) @@ -45,9 +43,9 @@ type SQLAdminClient interface { // GetDatabaseInstance returns database instance details for the project/instance // configured in a session. GetDatabaseInstance(ctx context.Context, db types.Database) (*sqladmin.DatabaseInstance, error) - // GenerateEphemeralCert returns a new client certificate with RSA key for the - // project/instance configured in a session. - GenerateEphemeralCert(ctx context.Context, db types.Database, certExpiry time.Time) (*tls.Certificate, error) + // GenerateEphemeralCert returns a new PEM-encoded client certificate for + // the project/instance configured in a session. + GenerateEphemeralCert(ctx context.Context, db types.Database, certExpiry time.Time, pubKey crypto.PublicKey) (string, error) } // NewGCPSQLAdminClient returns a GCPSQLAdminClient interface wrapping sqladmin.Service. @@ -99,42 +97,41 @@ func (g *gcpSQLAdminClient) GetDatabaseInstance(ctx context.Context, db types.Da return dbi, nil } -// GenerateEphemeralCert returns a new client certificate with RSA key created -// using the GenerateEphemeralCertRequest Cloud SQL API. Client certificates are -// required when enabling SSL in Cloud SQL. -func (g *gcpSQLAdminClient) GenerateEphemeralCert(ctx context.Context, db types.Database, certExpiry time.Time) (*tls.Certificate, error) { +// GenerateEphemeralCert returns a new client certificate created using the +// GenerateEphemeralCertRequest Cloud SQL API. Client certificates are required +// when enabling SSL in Cloud SQL. +func (g *gcpSQLAdminClient) GenerateEphemeralCert(ctx context.Context, db types.Database, certExpiry time.Time, pubKey crypto.PublicKey) (string, error) { // TODO(jimbishopp): cache database certificates to avoid expensive generate // operation on each connection. - // Generate RSA private key, x509 encoded public key, and append to certificate request. - pkey, err := rsa.GenerateKey(rand.Reader, constants.RSAKeySize) - if err != nil { - return nil, trace.Wrap(err) - } - pkix, err := x509.MarshalPKIXPublicKey(pkey.Public()) - if err != nil { - return nil, trace.Wrap(err) + var keyPEM []byte + switch pubKey.(type) { + case *rsa.PublicKey: + // keys.MarshalPublicKey would use PKCS#1 format for an RSA public key, + // we specifically want PKIX here. + pkix, err := x509.MarshalPKIXPublicKey(pubKey) + if err != nil { + return "", trace.Wrap(err) + } + keyPEM = pem.EncodeToMemory(&pem.Block{Bytes: pkix, Type: "RSA PUBLIC KEY"}) + default: + var err error + keyPEM, err = keys.MarshalPublicKey(pubKey) + if err != nil { + return "", trace.Wrap(err) + } } // Make API call. gcp := db.GetGCP() req := g.service.Connect.GenerateEphemeralCert(gcp.ProjectID, gcp.InstanceID, &sqladmin.GenerateEphemeralCertRequest{ - PublicKey: string(pem.EncodeToMemory(&pem.Block{Bytes: pkix, Type: "RSA PUBLIC KEY"})), + PublicKey: string(keyPEM), ValidDuration: fmt.Sprintf("%ds", int(time.Until(certExpiry).Seconds())), }) resp, err := req.Context(ctx).Do() if err != nil { - return nil, trace.Wrap(convertAPIError(err)) + return "", trace.Wrap(convertAPIError(err)) } - // Create TLS certificate from returned ephemeral certificate and private key. - keyPEM, err := keys.MarshalPrivateKey(pkey) - if err != nil { - return nil, trace.Wrap(err) - } - cert, err := tls.X509KeyPair([]byte(resp.EphemeralCert.Cert), keyPEM) - if err != nil { - return nil, trace.Wrap(err) - } - return &cert, nil + return resp.EphemeralCert.Cert, nil } diff --git a/lib/cloud/mocks/gcp.go b/lib/cloud/mocks/gcp.go index 8cf686d3d90..7981a515acc 100644 --- a/lib/cloud/mocks/gcp.go +++ b/lib/cloud/mocks/gcp.go @@ -20,7 +20,7 @@ package mocks import ( "context" - "crypto/tls" + "crypto" "time" "github.com/gravitational/trace" @@ -39,7 +39,7 @@ type GCPSQLAdminClientMock struct { // DatabaseInstance is returned from GetDatabaseInstance. DatabaseInstance *sqladmin.DatabaseInstance // EphemeralCert is returned from GenerateEphemeralCert. - EphemeralCert *tls.Certificate + EphemeralCert string // DatabaseUser is returned from GetUser. DatabaseUser *sqladmin.User } @@ -59,7 +59,7 @@ func (g *GCPSQLAdminClientMock) GetDatabaseInstance(ctx context.Context, db type return g.DatabaseInstance, nil } -func (g *GCPSQLAdminClientMock) GenerateEphemeralCert(ctx context.Context, db types.Database, certExpiry time.Time) (*tls.Certificate, error) { +func (g *GCPSQLAdminClientMock) GenerateEphemeralCert(_ context.Context, _ types.Database, _ time.Time, _ crypto.PublicKey) (string, error) { return g.EphemeralCert, nil } diff --git a/lib/srv/db/access_test.go b/lib/srv/db/access_test.go index 4266ab3361e..6069c1ae1cb 100644 --- a/lib/srv/db/access_test.go +++ b/lib/srv/db/access_test.go @@ -23,6 +23,7 @@ import ( "context" "crypto/tls" "database/sql" + "encoding/pem" "errors" "fmt" "io" @@ -759,6 +760,10 @@ func TestGCPRequireSSL(t *testing.T) { Username: user, }) require.NoError(t, err) + certPEM := pem.EncodeToMemory(&pem.Block{ + Type: "CERTIFICATE", + Bytes: ephemeralCert.Certificate[0], + }) // Setup database servers for Postgres and MySQL with a mock GCP API that // will require SSL and return the ephemeral certificate created above. @@ -768,7 +773,7 @@ func TestGCPRequireSSL(t *testing.T) { withCloudSQLMySQLTLS("mysql", user, cloudSQLPassword)(t, ctx, testCtx), }, GCPSQL: &mocks.GCPSQLAdminClientMock{ - EphemeralCert: ephemeralCert, + EphemeralCert: string(certPEM), DatabaseInstance: &sqladmin.DatabaseInstance{ Settings: &sqladmin.Settings{ IpConfiguration: &sqladmin.IpConfiguration{ diff --git a/lib/srv/db/auth_test.go b/lib/srv/db/auth_test.go index ef6860c7a8b..d58b4314832 100644 --- a/lib/srv/db/auth_test.go +++ b/lib/srv/db/auth_test.go @@ -35,8 +35,10 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/lib/cloud/mocks" "github.com/gravitational/teleport/lib/defaults" + "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/srv/db/common" ) @@ -404,6 +406,11 @@ func (a *testAuth) GetAWSIAMCreds(ctx context.Context, database types.Database, return atlasAuthUser, atlasAuthToken, atlasAuthSessionToken, nil } +func (a *testAuth) GenerateDatabaseClientKey(ctx context.Context) (*keys.PrivateKey, error) { + key, err := keys.ParsePrivateKey(fixtures.PEMBytes["rsa"]) + return key, trace.Wrap(err) +} + func (a *testAuth) WithLogger(getUpdatedLogger func(logrus.FieldLogger) logrus.FieldLogger) common.Auth { // TODO(greedy52) update WithLogger to use slog. return &testAuth{ diff --git a/lib/srv/db/cloud/gcp.go b/lib/srv/db/cloud/gcp.go index 91d5a7b3a7a..892b713e43f 100644 --- a/lib/srv/db/cloud/gcp.go +++ b/lib/srv/db/cloud/gcp.go @@ -26,6 +26,7 @@ import ( "github.com/gravitational/trace" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/lib/cloud/gcp" "github.com/gravitational/teleport/lib/srv/db/common" ) @@ -52,11 +53,25 @@ or "cloudsql.instances.get" IAM permission.`, err) return dbi.Settings.IpConfiguration.RequireSsl, nil } +// AppendGPCClientCertRequest is a request to update [TLSConfig] with an +// ephemeral GCP client certificate. +type AppendGCPClientCertRequest struct { + GCPClient gcp.SQLAdminClient + GenerateKey func(context.Context) (*keys.PrivateKey, error) + Expiry time.Time + Database types.Database + TLSConfig *tls.Config +} + // AppendGCPClientCert calls the GCP API to generate an ephemeral certificate // and adds it to the TLS config. An access denied error is returned when the // generate call fails. -func AppendGCPClientCert(ctx context.Context, certExpiry time.Time, database types.Database, gcpClient gcp.SQLAdminClient, tlsConfig *tls.Config) error { - cert, err := gcpClient.GenerateEphemeralCert(ctx, database, certExpiry) +func AppendGCPClientCert(ctx context.Context, req *AppendGCPClientCertRequest) error { + privateKey, err := req.GenerateKey(ctx) + if err != nil { + return trace.Wrap(err) + } + certPEM, err := req.GCPClient.GenerateEphemeralCert(ctx, req.Database, req.Expiry, privateKey.Public()) if err != nil { err = common.ConvertError(err) if trace.IsAccessDenied(err) { @@ -67,8 +82,12 @@ func AppendGCPClientCert(ctx context.Context, certExpiry time.Time, database typ Make sure Teleport db service has "Cloud SQL Admin" GCP IAM role, or "cloudsql.sslCerts.createEphemeral" IAM permission.`, err) } - return trace.Wrap(err, "Failed to generate GCP ephemeral client certificate for %q.", database.GetGCP().GetServerName()) + return trace.Wrap(err, "Failed to generate GCP ephemeral client certificate for %q.", req.Database.GetGCP().GetServerName()) } - tlsConfig.Certificates = []tls.Certificate{*cert} + tlsCert, err := privateKey.TLSCertificate([]byte(certPEM)) + if err != nil { + return trace.Wrap(err) + } + req.TLSConfig.Certificates = []tls.Certificate{tlsCert} return nil } diff --git a/lib/srv/db/common/auth.go b/lib/srv/db/common/auth.go index 4e0adc88975..5cee3d170d9 100644 --- a/lib/srv/db/common/auth.go +++ b/lib/srv/db/common/auth.go @@ -51,12 +51,13 @@ import ( "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/types" azureutils "github.com/gravitational/teleport/api/utils/azure" + "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/api/utils/retryutils" - "github.com/gravitational/teleport/lib/auth/native" "github.com/gravitational/teleport/lib/cloud" awslib "github.com/gravitational/teleport/lib/cloud/aws" libazure "github.com/gravitational/teleport/lib/cloud/azure" "github.com/gravitational/teleport/lib/cloud/gcp" + "github.com/gravitational/teleport/lib/cryptosuites" "github.com/gravitational/teleport/lib/defaults" dbiam "github.com/gravitational/teleport/lib/srv/db/common/iam" "github.com/gravitational/teleport/lib/tlsca" @@ -101,6 +102,9 @@ type Auth interface { // GetAWSIAMCreds returns the AWS IAM credentials, including access key, // secret access key and session token. GetAWSIAMCreds(ctx context.Context, database types.Database, databaseUser string) (string, string, string, error) + // GenerateDatabaseClientKey generates a cryptographic key appropriate for + // database client connections. + GenerateDatabaseClientKey(context.Context) (*keys.PrivateKey, error) // WithLogger returns a new instance of Auth with updated logger. // The callback function receives the current logger and returns a new one. WithLogger(getUpdatedLogger func(logrus.FieldLogger) logrus.FieldLogger) Auth @@ -978,7 +982,7 @@ func verifyConnectionFunc(rootCAs *x509.CertPool) func(cs tls.ConnectionState) e // getClientCert signs an ephemeral client certificate used by this // server to authenticate with the database instance. func (a *dbAuth) getClientCert(ctx context.Context, expiry time.Time, databaseUser string) (cert *tls.Certificate, cas [][]byte, err error) { - privateKey, err := native.GeneratePrivateKey() + privateKey, err := a.GenerateDatabaseClientKey(ctx) if err != nil { return nil, nil, trace.Wrap(err) } @@ -1009,6 +1013,21 @@ func (a *dbAuth) getClientCert(ctx context.Context, expiry time.Time, databaseUs return &clientCert, resp.CACerts, nil } +// GenerateDatabaseClientKey generates a cryptographic key appropriate for +// database client connections. +func (a *dbAuth) GenerateDatabaseClientKey(ctx context.Context) (*keys.PrivateKey, error) { + signer, err := cryptosuites.GenerateKey(ctx, + cryptosuites.GetCurrentSuiteFromAuthPreference(a), cryptosuites.DatabaseClient) + if err != nil { + return nil, trace.Wrap(err) + } + privateKey, err := keys.NewSoftwarePrivateKey(signer) + if err != nil { + return nil, trace.Wrap(err) + } + return privateKey, nil +} + // GetAuthPreference returns the cluster authentication config. func (a *dbAuth) GetAuthPreference(ctx context.Context) (types.AuthPreference, error) { return a.cfg.AuthClient.GetAuthPreference(ctx) diff --git a/lib/srv/db/common/test.go b/lib/srv/db/common/test.go index 65b6cddb9a6..c1be10cb3e8 100644 --- a/lib/srv/db/common/test.go +++ b/lib/srv/db/common/test.go @@ -33,7 +33,6 @@ import ( "github.com/gravitational/teleport/api/utils/keys" "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/authclient" - "github.com/gravitational/teleport/lib/auth/testauthority" "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/tlsca" @@ -123,7 +122,7 @@ func MakeTestServerTLSConfig(config TestServerConfig) (*tls.Config, error) { if cn == "" { cn = "localhost" } - privateKey, err := testauthority.New().GeneratePrivateKey() + privateKey, err := keys.ParsePrivateKey(fixtures.PEMBytes["rsa"]) if err != nil { return nil, trace.Wrap(err) } diff --git a/lib/srv/db/mysql/engine.go b/lib/srv/db/mysql/engine.go index 4355c836f3f..5ba7af2294d 100644 --- a/lib/srv/db/mysql/engine.go +++ b/lib/srv/db/mysql/engine.go @@ -256,7 +256,13 @@ func (e *Engine) connect(ctx context.Context, sessionCtx *common.Session) (*clie // the instance requires SSL. Also use a TLS dialer instead of // the default net dialer when GCP requires SSL. if requireSSL { - err = cloud.AppendGCPClientCert(ctx, sessionCtx.GetExpiry(), sessionCtx.Database, gcpClient, tlsConfig) + err = cloud.AppendGCPClientCert(ctx, &cloud.AppendGCPClientCertRequest{ + GCPClient: gcpClient, + GenerateKey: e.Auth.GenerateDatabaseClientKey, + Expiry: sessionCtx.GetExpiry(), + Database: sessionCtx.Database, + TLSConfig: tlsConfig, + }) if err != nil { return nil, trace.Wrap(err) } diff --git a/lib/srv/db/postgres/connector.go b/lib/srv/db/postgres/connector.go index 08bdae1fb3a..81873b6afd7 100644 --- a/lib/srv/db/postgres/connector.go +++ b/lib/srv/db/postgres/connector.go @@ -104,7 +104,13 @@ func (c *connector) getConnectConfig(ctx context.Context) (*pgconn.Config, error // Create ephemeral certificate and append to TLS config when // the instance requires SSL. if requireSSL { - err = cloud.AppendGCPClientCert(ctx, c.certExpiry, c.database, gcpClient, config.TLSConfig) + err = cloud.AppendGCPClientCert(ctx, &cloud.AppendGCPClientCertRequest{ + GCPClient: gcpClient, + GenerateKey: c.auth.GenerateDatabaseClientKey, + Expiry: c.certExpiry, + Database: c.database, + TLSConfig: config.TLSConfig, + }) if err != nil { return nil, trace.Wrap(err) }