From 96ed71b2f7199f9ddc12b3ace6dee505eada5ddf Mon Sep 17 00:00:00 2001 From: Gabriel Corado Date: Wed, 29 Nov 2023 09:50:00 -0300 Subject: [PATCH] Add token and database URI validation to SQL Server configure script route (#34035) * feat(web): add token and database uri validations to sql server configure * feat: add sql server address validation * test: update example sql server URIs * Update lib/web/databases_test.go Co-authored-by: STeve (Xin) Huang --------- Co-authored-by: STeve (Xin) Huang --- lib/config/configuration_test.go | 4 +- lib/service/servicecfg/config_test.go | 2 +- lib/services/database.go | 52 +++++++++++++ lib/services/database_test.go | 22 ++++++ lib/web/databases.go | 22 ++++-- lib/web/databases_test.go | 78 ++++++++++++++++++- lib/web/join_tokens.go | 20 +++-- .../database/sqlserver/configure-ad.ps1 | 16 ++-- 8 files changed, 189 insertions(+), 27 deletions(-) diff --git a/lib/config/configuration_test.go b/lib/config/configuration_test.go index aa09319dcf6..ab89252c102 100644 --- a/lib/config/configuration_test.go +++ b/lib/config/configuration_test.go @@ -2784,7 +2784,7 @@ func TestDatabaseCLIFlags(t *testing.T) { inFlags: CommandLineFlags{ DatabaseName: "sqlserver", DatabaseProtocol: defaults.ProtocolSQLServer, - DatabaseURI: "localhost:1433", + DatabaseURI: "sqlserver.example.com:1433", DatabaseADKeytabFile: "/etc/keytab", DatabaseADDomain: "EXAMPLE.COM", DatabaseADSPN: "MSSQLSvc/sqlserver.example.com:1433", @@ -2792,7 +2792,7 @@ func TestDatabaseCLIFlags(t *testing.T) { outDatabase: servicecfg.Database{ Name: "sqlserver", Protocol: defaults.ProtocolSQLServer, - URI: "localhost:1433", + URI: "sqlserver.example.com:1433", TLS: servicecfg.DatabaseTLS{ Mode: servicecfg.VerifyFull, }, diff --git a/lib/service/servicecfg/config_test.go b/lib/service/servicecfg/config_test.go index 325855d8b45..d10e3780187 100644 --- a/lib/service/servicecfg/config_test.go +++ b/lib/service/servicecfg/config_test.go @@ -259,7 +259,7 @@ func TestCheckDatabase(t *testing.T) { inDatabase: Database{ Name: "sqlserver", Protocol: defaults.ProtocolSQLServer, - URI: "localhost:1433", + URI: "sqlserver.example.com:1433", AD: DatabaseAD{ KeytabFile: "/etc/keytab", Domain: "test-domain", diff --git a/lib/services/database.go b/lib/services/database.go index f538ac3796e..fbbe89179ba 100644 --- a/lib/services/database.go +++ b/lib/services/database.go @@ -21,6 +21,7 @@ import ( "errors" "fmt" "net" + "net/netip" "net/url" "strings" @@ -176,6 +177,10 @@ func ValidateDatabase(db types.Database) error { if err := validateClickhouseURI(db); err != nil { return trace.Wrap(err) } + } else if db.GetProtocol() == defaults.ProtocolSQLServer { + if err := ValidateSQLServerURI(db.GetURI()); err != nil { + return trace.BadParameter("invalid SQL Server address: %v", err) + } } else if needsURIValidation(db) { if _, _, err := net.SplitHostPort(db.GetURI()); err != nil { return trace.BadParameter("invalid database %q address %q: %v", db.GetName(), db.GetURI(), err) @@ -310,6 +315,51 @@ func validateMongoDB(db types.Database) error { return nil } +// ValidateSQLServerURI validates SQL Server URI and returns host and +// port. +// +// Since Teleport only supports SQL Server authentcation using AD (self-hosted +// or Azure) the database URI must include: computer name, domain and port. +// +// A few examples of valid URIs: +// - computer.ad.example.com:1433 +// - computer.domain.com:1433 +func ValidateSQLServerURI(uri string) error { + // Add a temporary schema to make a valid URL for url.Parse if schema is + // not found. + if !strings.Contains(uri, "://") { + uri = sqlServerSchema + "://" + uri + } + + parsedURI, err := url.Parse(uri) + if err != nil { + return trace.BadParameter("unabled to parse database address: %s", err) + } + + if parsedURI.Scheme != sqlServerSchema { + return trace.BadParameter("only %q is supported as database address schema", sqlServerSchema) + } + + if parsedURI.Port() == "" { + return trace.BadParameter("database address must include port") + } + + if parsedURI.Path != "" { + return trace.BadParameter("database address with database name is not supported") + } + + if _, err := netip.ParseAddr(parsedURI.Hostname()); err == nil { + return trace.BadParameter("database address as IP is not supported, use URI with domain and computer name instead") + } + + parts := strings.Split(parsedURI.Hostname(), ".") + if len(parts) < 3 { + return trace.BadParameter("database address must include domain and computer name") + } + + return nil +} + func isDNSError(err error) bool { if err == nil { return false @@ -2013,4 +2063,6 @@ const ( const ( // azureSQLServerDefaultPort is the default port for Azure SQL Server. azureSQLServerDefaultPort = 1433 + // sqlServerSchema is the SQL Server schema. + sqlServerSchema = "mssql" ) diff --git a/lib/services/database_test.go b/lib/services/database_test.go index 9db56c370bb..a94f08ed63e 100644 --- a/lib/services/database_test.go +++ b/lib/services/database_test.go @@ -483,6 +483,28 @@ func TestValidateDatabase(t *testing.T) { } } +func TestValidateSQLServerDatabaseURI(t *testing.T) { + for _, test := range []struct { + uri string + assertErr require.ErrorAssertionFunc + }{ + {"mssql://computer.domain.com:1433", require.NoError}, + {"computer.domain.com:1433", require.NoError}, + {"computer.ad.domain.com:1433", require.NoError}, + {"computer.ad.domain.com:1433/hello", require.Error}, + {"mssql://computer.domain.com:1433/hello", require.Error}, + {"computer.domain.com", require.Error}, + {"computer.com:1433", require.Error}, + {"0.0.0.0:1433", require.Error}, + {"mssql://", require.Error}, + {"http://computer.domain.com:1433", require.Error}, + } { + t.Run(test.uri, func(t *testing.T) { + test.assertErr(t, ValidateSQLServerURI(test.uri)) + }) + } +} + // indent returns the string where each line is indented by the specified // number of spaces. func indent(s string, spaces int) string { diff --git a/lib/web/databases.go b/lib/web/databases.go index c030f6f4e5f..1d35d01703c 100644 --- a/lib/web/databases.go +++ b/lib/web/databases.go @@ -23,7 +23,9 @@ import ( "encoding/json" "encoding/pem" "fmt" + "net" "net/http" + "net/url" "github.com/gravitational/trace" "github.com/julienschmidt/httprouter" @@ -34,6 +36,7 @@ import ( "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/httplib" "github.com/gravitational/teleport/lib/reversetunnelclient" + "github.com/gravitational/teleport/lib/services" dbiam "github.com/gravitational/teleport/lib/srv/db/common/iam" "github.com/gravitational/teleport/lib/web/scripts" "github.com/gravitational/teleport/lib/web/ui" @@ -294,18 +297,17 @@ func (h *Handler) handleDatabaseGetIAMPolicy(w http.ResponseWriter, r *http.Requ func (h *Handler) sqlServerConfigureADScriptHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (interface{}, error) { tokenStr := p.ByName("token") - if tokenStr == "" { - return "", trace.BadParameter("invalid token") + if err := validateJoinToken(tokenStr); err != nil { + return "", trace.Wrap(err) } dbAddress := r.URL.Query().Get("uri") - if dbAddress == "" { - return "", trace.BadParameter("invalid database address") + if err := services.ValidateSQLServerURI(dbAddress); err != nil { + return "", trace.BadParameter("invalid database address: %v", err) } // verify that the token exists - _, err := h.GetProxyClient().GetToken(r.Context(), tokenStr) - if err != nil { + if _, err := h.GetProxyClient().GetToken(r.Context(), tokenStr); err != nil { return "", trace.BadParameter("invalid token") } @@ -347,6 +349,12 @@ func (h *Handler) sqlServerConfigureADScriptHandle(w http.ResponseWriter, r *htt return nil, trace.BadParameter("no PEM data in CA data") } + // Split host and port so we can escape domain characters. + dbHost, dbPort, err := net.SplitHostPort(dbAddress) + if err != nil { + return nil, trace.Wrap(err) + } + httplib.SetScriptHeaders(w.Header()) w.WriteHeader(http.StatusOK) err = scripts.DatabaseAccessSQLServerConfigureScript.Execute(w, scripts.DatabaseAccessSQLServerConfigureParams{ @@ -356,7 +364,7 @@ func (h *Handler) sqlServerConfigureADScriptHandle(w http.ResponseWriter, r *htt CRLPEM: string(encodeCRLPEM(caCRL)), ProxyPublicAddr: proxyServers[0].GetPublicAddr(), ProvisionToken: tokenStr, - DBAddress: dbAddress, + DBAddress: net.JoinHostPort(url.QueryEscape(dbHost), dbPort), }) return nil, trace.Wrap(err) diff --git a/lib/web/databases_test.go b/lib/web/databases_test.go index 6d16ff13947..9cc140265e3 100644 --- a/lib/web/databases_test.go +++ b/lib/web/databases_test.go @@ -19,8 +19,10 @@ package web import ( "context" "encoding/json" + "fmt" "net/http" "net/url" + "regexp" "testing" "time" @@ -414,7 +416,7 @@ func TestHandleSQLServerConfigureScript(t *testing.T) { }{ { desc: "valid token and uri", - uri: "instance.example.teleport.dev", + uri: "instance.example.teleport.dev:1433", tokenFunc: func(t *testing.T) string { pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) @@ -423,7 +425,7 @@ func TestHandleSQLServerConfigureScript(t *testing.T) { assertError: require.NoError, }, { - desc: "valid token and invalid uri", + desc: "valid token and empty uri", uri: "", tokenFunc: func(t *testing.T) string { pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) @@ -432,9 +434,49 @@ func TestHandleSQLServerConfigureScript(t *testing.T) { }, assertError: require.Error, }, + { + desc: "valid token and invalid uri", + uri: "hello#hello", + tokenFunc: func(t *testing.T) string { + pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) + require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) + return token + }, + assertError: require.Error, + }, + { + desc: "invalid line break character token and invalid uri", + uri: "computer.domain\n.com:1433", + tokenFunc: func(t *testing.T) string { + pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) + require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) + return token + }, + assertError: require.Error, + }, + { + desc: "invalid character ` token and invalid uri", + uri: "computer.domain`.com:1433", + tokenFunc: func(t *testing.T) string { + pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) + require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) + return token + }, + assertError: require.Error, + }, + { + desc: "invalid character | token and invalid uri", + uri: "computer.domain|.com:1433", + tokenFunc: func(t *testing.T) string { + pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) + require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) + return token + }, + assertError: require.Error, + }, { desc: "invalid token", - uri: "instance.example.teleport.dev", + uri: "instance.example.teleport.dev:1433", tokenFunc: func(_ *testing.T) string { return "random-token" }, assertError: require.Error, }, @@ -448,6 +490,36 @@ func TestHandleSQLServerConfigureScript(t *testing.T) { tc.assertError(t, err) }) } + +} + +// TestHandleSQLServerConfigureScriptDatabaseURIEscaped given a SQL Server +// database URI, ensures that special characters are escaped when placed on the +// PowerShell script. +func TestHandleSQLServerConfigureScriptDatabaseURIEscaped(t *testing.T) { + ctx := context.Background() + env := newWebPack(t, 1) + proxy := env.proxies[0] + pack := proxy.authPack(t, "user", nil /* roles */) + pt, token := generateProvisionToken(t, types.RoleDatabase, env.clock.Now().Add(time.Hour)) + require.NoError(t, env.server.Auth().CreateToken(ctx, pt)) + re := regexp.MustCompile(`\$DB_ADDRESS\s*=\s*'([^']+)'`) + + for _, c := range []string{";", "\"", "'", "&", "$", "(", ")"} { + t.Run(c, func(t *testing.T) { + uri := fmt.Sprintf("database.ad%s.com:1433", c) + resp, err := pack.clt.Get( + ctx, + pack.clt.Endpoint("webapi/scripts/databases/configure/sqlserver", token, "configure-ad.ps1"), + url.Values{"uri": []string{uri}}, + ) + require.NoError(t, err) + escapedURIResult := re.FindStringSubmatch(string(resp.Bytes())) + require.Len(t, escapedURIResult, 2) + require.NotEqual(t, uri, escapedURIResult[1]) + require.Contains(t, escapedURIResult[1], url.QueryEscape(c)) + }) + } } func mustCreateDatabaseServer(t *testing.T, db *types.DatabaseV3) types.DatabaseServer { diff --git a/lib/web/join_tokens.go b/lib/web/join_tokens.go index 9423b24b662..03bf6a174e5 100644 --- a/lib/web/join_tokens.go +++ b/lib/web/join_tokens.go @@ -302,14 +302,9 @@ func (h *Handler) getDatabaseJoinScriptHandle(w http.ResponseWriter, r *http.Req func getJoinScript(ctx context.Context, settings scriptSettings, m nodeAPIGetter) (string, error) { switch types.JoinMethod(settings.joinMethod) { case types.JoinMethodUnspecified, types.JoinMethodToken: - decodedToken, err := hex.DecodeString(settings.token) - if err != nil { + if err := validateJoinToken(settings.token); err != nil { return "", trace.Wrap(err) } - if len(decodedToken) != auth.TokenLenBytes { - return "", trace.BadParameter("invalid token %q", decodedToken) - } - case types.JoinMethodIAM: default: return "", trace.BadParameter("join method %q is not supported via script", settings.joinMethod) @@ -436,6 +431,19 @@ func getJoinScript(ctx context.Context, settings scriptSettings, m nodeAPIGetter return buf.String(), nil } +// validateJoinToken validate a join token. +func validateJoinToken(token string) error { + decodedToken, err := hex.DecodeString(token) + if err != nil { + return trace.BadParameter("invalid token %q", token) + } + if len(decodedToken) != auth.TokenLenBytes { + return trace.BadParameter("invalid token %q", decodedToken) + } + + return nil +} + // generateIAMTokenName makes a deterministic name for a iam join token // based on its rule set func generateIAMTokenName(rules []*types.TokenRule) (string, error) { diff --git a/lib/web/scripts/database/sqlserver/configure-ad.ps1 b/lib/web/scripts/database/sqlserver/configure-ad.ps1 index 04b65bd13e2..451061cc302 100644 --- a/lib/web/scripts/database/sqlserver/configure-ad.ps1 +++ b/lib/web/scripts/database/sqlserver/configure-ad.ps1 @@ -1,13 +1,13 @@ $ErrorActionPreference = "Stop" -$TELEPORT_CA_CERT_PEM = "{{.CACertPEM}}" -$TELEPORT_CA_CERT_SHA1 = "{{.CACertSHA1}}" -$TELEPORT_CA_CERT_BLOB_BASE64 = "{{.CACertBase64}}" -$TELEPORT_CRL_PEM = "{{.CRLPEM}}" -$TELEPORT_PROXY_PUBLIC_ADDR = "{{.ProxyPublicAddr}}" -$TELEPORT_PROVISION_TOKEN = "{{.ProvisionToken}}" +$TELEPORT_CA_CERT_PEM = '{{.CACertPEM}}' +$TELEPORT_CA_CERT_SHA1 = '{{.CACertSHA1}}' +$TELEPORT_CA_CERT_BLOB_BASE64 = '{{.CACertBase64}}' +$TELEPORT_CRL_PEM = '{{.CRLPEM}}' +$TELEPORT_PROXY_PUBLIC_ADDR = '{{.ProxyPublicAddr}}' +$TELEPORT_PROVISION_TOKEN = '{{.ProvisionToken}}' -$DB_ADDRESS = "{{.DBAddress}}" +$DB_ADDRESS = '{{.DBAddress}}' $COMPUTER_NAME = ($DB_ADDRESS -split '\.')[0].ToLower() $DOMAIN_NAME=(Get-ADDomain).DNSRoot @@ -128,4 +128,4 @@ Write-Output $OUTPUT Remove-Item $TeleportPEMFile -Recurse Remove-Item $TeleportCRLFile -Recurse Remove-Item $WindowsDERFile -Recurse -Remove-Item $WindowsPEMFile -Recurse \ No newline at end of file +Remove-Item $WindowsPEMFile -Recurse