From eadbe2d8a1e2dfb8730c37a2fae4e7091da0682a Mon Sep 17 00:00:00 2001 From: Hugo Shaka Date: Fri, 21 Feb 2025 13:44:03 -0500 Subject: [PATCH] Refactor node-join script to take safer options and reuse install option logic (#52196) * Add install script using teleport-update and oneoff.sh * Refactor node-join script to take safer options and reuse install option logic * GoDoc + make functions private * Address edoardo's feedback --- lib/web/apiserver.go | 2 +- lib/web/apiserver_test.go | 1 + lib/web/join_tokens.go | 250 +++------------ lib/web/join_tokens_test.go | 444 ++++++++++++++++----------- lib/web/scripts/install_node.go | 164 +++++++++- lib/web/scripts/install_node_test.go | 2 +- 6 files changed, 473 insertions(+), 390 deletions(-) diff --git a/lib/web/apiserver.go b/lib/web/apiserver.go index 009083aa9de..e4318f78e52 100644 --- a/lib/web/apiserver.go +++ b/lib/web/apiserver.go @@ -2250,7 +2250,7 @@ func (h *Handler) installer(w http.ResponseWriter, r *http.Request, p httprouter // https://updates.releases.teleport.dev/v1/stable/cloud/version installUpdater := automaticUpgrades(*ping.ServerFeatures) if installUpdater { - repoChannel = stableCloudChannelRepo + repoChannel = automaticupgrades.DefaultCloudChannelName } azureClientID := r.URL.Query().Get("azure-client-id") diff --git a/lib/web/apiserver_test.go b/lib/web/apiserver_test.go index a9d93db9559..ced37e72729 100644 --- a/lib/web/apiserver_test.go +++ b/lib/web/apiserver_test.go @@ -3643,6 +3643,7 @@ func TestKnownWebPathsWithAndWithoutV1Prefix(t *testing.T) { func TestInstallDatabaseScriptGeneration(t *testing.T) { const username = "test-user@example.com" + modules.SetTestModules(t, &modules.TestModules{TestBuildType: modules.BuildCommunity}) // Users should be able to create Tokens even if they can't update them roleTokenCRD, err := types.NewRole(services.RoleNameForUser(username), types.RoleSpecV6{ diff --git a/lib/web/join_tokens.go b/lib/web/join_tokens.go index 9a1929f6061..511e2b71963 100644 --- a/lib/web/join_tokens.go +++ b/lib/web/join_tokens.go @@ -19,7 +19,6 @@ package web import ( - "bytes" "context" "encoding/hex" "fmt" @@ -27,25 +26,19 @@ import ( "net/http" "net/url" "reflect" - "regexp" "sort" - "strconv" "strings" "time" - "github.com/google/safetext/shsprintf" "github.com/google/uuid" "github.com/gravitational/trace" "github.com/julienschmidt/httprouter" - "k8s.io/apimachinery/pkg/util/validation" - "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/client/proto" "github.com/gravitational/teleport/api/types" apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/httplib" - "github.com/gravitational/teleport/lib/modules" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/ui" @@ -55,8 +48,7 @@ import ( ) const ( - stableCloudChannelRepo = "stable/cloud" - HeaderTokenName = "X-Teleport-TokenName" + HeaderTokenName = "X-Teleport-TokenName" ) // nodeJoinToken contains node token fields for the UI. @@ -80,15 +72,9 @@ type scriptSettings struct { appURI string joinMethod string databaseInstallMode bool - installUpdater bool discoveryInstallMode bool discoveryGroup string - - // automaticUpgradesVersion is the target automatic upgrades version. - // The version must be valid semver, with the leading 'v'. e.g. v15.0.0-dev - // Required when installUpdater is true. - automaticUpgradesVersion string } // automaticUpgrades returns whether automaticUpgrades should be enabled. @@ -377,41 +363,16 @@ func (h *Handler) createTokenForDiscoveryHandle(w http.ResponseWriter, r *http.R }, nil } -// getAutoUpgrades checks if automaticUpgrades are enabled and returns the -// version that should be used according to auto upgrades default channel. -// If something bad happens, the error is logged and the function falls back to -// the process Teleport version. -func (h *Handler) getAutoUpgrades(ctx context.Context) (bool, string) { - var autoUpgradesVersion string - var err error - autoUpgrades := automaticUpgrades(h.GetClusterFeatures()) - if autoUpgrades { - const group, updaterUUID = "", "" - autoUpgradesVersion, err = h.autoUpdateAgentVersion(ctx, group, updaterUUID) - if err != nil { - h.logger.WarnContext(ctx, "Failed to get auto upgrades version, falling back to self version.", "error", err) - return autoUpgrades, teleport.Version - } - autoUpgradesVersion = fmt.Sprintf("v%s", autoUpgradesVersion) - } - return autoUpgrades, autoUpgradesVersion - -} - func (h *Handler) getNodeJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (interface{}, error) { httplib.SetScriptHeaders(w.Header()) - autoUpgrades, autoUpgradesVersion := h.getAutoUpgrades(r.Context()) - settings := scriptSettings{ - token: params.ByName("token"), - appInstallMode: false, - joinMethod: r.URL.Query().Get("method"), - installUpdater: autoUpgrades, - automaticUpgradesVersion: autoUpgradesVersion, + token: params.ByName("token"), + appInstallMode: false, + joinMethod: r.URL.Query().Get("method"), } - script, err := getJoinScript(r.Context(), settings, h.GetProxyClient()) + script, err := h.getJoinScript(r.Context(), settings) if err != nil { h.logger.InfoContext(r.Context(), "Failed to return the node install script", "error", err) w.Write(scripts.ErrorBashScript) @@ -451,18 +412,14 @@ func (h *Handler) getAppJoinScriptHandle(w http.ResponseWriter, r *http.Request, return nil, nil } - autoUpgrades, autoUpgradesVersion := h.getAutoUpgrades(r.Context()) - settings := scriptSettings{ - token: params.ByName("token"), - appInstallMode: true, - appName: name, - appURI: uri, - installUpdater: autoUpgrades, - automaticUpgradesVersion: autoUpgradesVersion, + token: params.ByName("token"), + appInstallMode: true, + appName: name, + appURI: uri, } - script, err := getJoinScript(r.Context(), settings, h.GetProxyClient()) + script, err := h.getJoinScript(r.Context(), settings) if err != nil { h.logger.InfoContext(r.Context(), "Failed to return the app install script", "error", err) w.Write(scripts.ErrorBashScript) @@ -481,16 +438,12 @@ func (h *Handler) getAppJoinScriptHandle(w http.ResponseWriter, r *http.Request, func (h *Handler) getDatabaseJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (interface{}, error) { httplib.SetScriptHeaders(w.Header()) - autoUpgrades, autoUpgradesVersion := h.getAutoUpgrades(r.Context()) - settings := scriptSettings{ - token: params.ByName("token"), - databaseInstallMode: true, - installUpdater: autoUpgrades, - automaticUpgradesVersion: autoUpgradesVersion, + token: params.ByName("token"), + databaseInstallMode: true, } - script, err := getJoinScript(r.Context(), settings, h.GetProxyClient()) + script, err := h.getJoinScript(r.Context(), settings) if err != nil { h.logger.InfoContext(r.Context(), "Failed to return the database install script", "error", err) w.Write(scripts.ErrorBashScript) @@ -511,8 +464,6 @@ func (h *Handler) getDiscoveryJoinScriptHandle(w http.ResponseWriter, r *http.Re queryValues := r.URL.Query() const discoveryGroupQueryParam = "discoveryGroup" - autoUpgrades, autoUpgradesVersion := h.getAutoUpgrades(r.Context()) - discoveryGroup, err := url.QueryUnescape(queryValues.Get(discoveryGroupQueryParam)) if err != nil { h.logger.DebugContext(r.Context(), "Failed to return the discovery install script", @@ -531,14 +482,12 @@ func (h *Handler) getDiscoveryJoinScriptHandle(w http.ResponseWriter, r *http.Re } settings := scriptSettings{ - token: params.ByName("token"), - discoveryInstallMode: true, - discoveryGroup: discoveryGroup, - installUpdater: autoUpgrades, - automaticUpgradesVersion: autoUpgradesVersion, + token: params.ByName("token"), + discoveryInstallMode: true, + discoveryGroup: discoveryGroup, } - script, err := getJoinScript(r.Context(), settings, h.GetProxyClient()) + script, err := h.getJoinScript(r.Context(), settings) if err != nil { h.logger.InfoContext(r.Context(), "Failed to return the discovery install script", "error", err) w.Write(scripts.ErrorBashScript) @@ -554,8 +503,9 @@ func (h *Handler) getDiscoveryJoinScriptHandle(w http.ResponseWriter, r *http.Re return nil, nil } -func getJoinScript(ctx context.Context, settings scriptSettings, m nodeAPIGetter) (string, error) { - switch types.JoinMethod(settings.joinMethod) { +func (h *Handler) getJoinScript(ctx context.Context, settings scriptSettings) (string, error) { + joinMethod := types.JoinMethod(settings.joinMethod) + switch joinMethod { case types.JoinMethodUnspecified, types.JoinMethodToken: if err := validateJoinToken(settings.token); err != nil { return "", trace.Wrap(err) @@ -565,141 +515,55 @@ func getJoinScript(ctx context.Context, settings scriptSettings, m nodeAPIGetter return "", trace.BadParameter("join method %q is not supported via script", settings.joinMethod) } + clt := h.GetProxyClient() + // The provided token can be attacker controlled, so we must validate // it with the backend before using it to generate the script. - token, err := m.GetToken(ctx, settings.token) + token, err := clt.GetToken(ctx, settings.token) if err != nil { return "", trace.BadParameter("invalid token") } - // Get hostname and port from proxy server address. - proxyServers, err := m.GetProxies() - if err != nil { - return "", trace.Wrap(err) - } - - if len(proxyServers) == 0 { - return "", trace.NotFound("no proxy servers found") - } - - version := proxyServers[0].GetTeleportVersion() - - publicAddr := proxyServers[0].GetPublicAddr() - if publicAddr == "" { - return "", trace.Errorf("proxy public_addr is not set, you must set proxy_service.public_addr to the publicly reachable address of the proxy before you can generate a node join script") - } - - hostname, portStr, err := utils.SplitHostPort(publicAddr) - if err != nil { - return "", trace.Wrap(err) - } + // TODO(hugoShaka): hit the local accesspoint which has a cache instead of asking the auth every time. // Get the CA pin hashes of the cluster to join. - localCAResponse, err := m.GetClusterCACert(ctx) + localCAResponse, err := clt.GetClusterCACert(ctx) if err != nil { return "", trace.Wrap(err) } + caPins, err := tlsca.CalculatePins(localCAResponse.TLSCA) if err != nil { return "", trace.Wrap(err) } - labelsList := []string{} - for labelKey, labelValues := range token.GetSuggestedLabels() { - labels := strings.Join(labelValues, " ") - labelsList = append(labelsList, fmt.Sprintf("%s=%s", labelKey, labels)) - } - - var dbServiceResourceLabels []string - if settings.databaseInstallMode { - suggestedAgentMatcherLabels := token.GetSuggestedAgentMatcherLabels() - dbServiceResourceLabels, err = scripts.MarshalLabelsYAML(suggestedAgentMatcherLabels, 6) - if err != nil { - return "", trace.Wrap(err) - } - } - - var buf bytes.Buffer - var appServerResourceLabels []string - // If app install mode is requested but parameters are blank for some reason, - // we need to return an error. - if settings.appInstallMode { - if errs := validation.IsDNS1035Label(settings.appName); len(errs) > 0 { - return "", trace.BadParameter("appName %q must be a valid DNS subdomain: https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#application-name", settings.appName) - } - if !appURIPattern.MatchString(settings.appURI) { - return "", trace.BadParameter("appURI %q contains invalid characters", settings.appURI) - } - - suggestedLabels := token.GetSuggestedLabels() - appServerResourceLabels, err = scripts.MarshalLabelsYAML(suggestedLabels, 4) - if err != nil { - return "", trace.Wrap(err) - } - } - - if settings.discoveryInstallMode { - if settings.discoveryGroup == "" { - return "", trace.BadParameter("discovery group is required") - } - } - - packageName := types.PackageNameOSS - if modules.GetModules().BuildType() == modules.BuildEnterprise { - packageName = types.PackageNameEnt - } - - // By default, it will use `stable/v`, eg stable/v12 - repoChannel := "" - - // The install script will install the updater (teleport-ent-updater) for Cloud customers enrolled in Automatic Upgrades. - // The repo channel used must be `stable/cloud` which has the available packages for the Cloud Customer's agents. - // It pins the teleport version to the one specified by the default version channel - // This ensures the initial installed version is the same as the `teleport-ent-updater` would install. - if settings.installUpdater { - if settings.automaticUpgradesVersion == "" { - return "", trace.Wrap(err, "automatic upgrades version must be set when installUpdater is true") - } - - repoChannel = stableCloudChannelRepo - // automaticUpgradesVersion has vX.Y.Z format, however the script - // expects the version to not include the `v` so we strip it - version = strings.TrimPrefix(settings.automaticUpgradesVersion, "v") - } - - // This section relies on Go's default zero values to make sure that the settings - // are correct when not installing an app. - err = scripts.InstallNodeBashScript.Execute(&buf, map[string]interface{}{ - "token": settings.token, - "hostname": hostname, - "port": portStr, - // The install.sh script has some manually generated configs and some - // generated by the `teleport config` commands. The old bash - // version used space delimited values whereas the teleport command uses - // a comma delimeter. The Old version can be removed when the install.sh - // file has been completely converted over. - "caPinsOld": strings.Join(caPins, " "), - "caPins": strings.Join(caPins, ","), - "packageName": packageName, - "repoChannel": repoChannel, - "installUpdater": strconv.FormatBool(settings.installUpdater), - "version": shsprintf.EscapeDefaultContext(version), - "appInstallMode": strconv.FormatBool(settings.appInstallMode), - "appServerResourceLabels": appServerResourceLabels, - "appName": shsprintf.EscapeDefaultContext(settings.appName), - "appURI": shsprintf.EscapeDefaultContext(settings.appURI), - "joinMethod": shsprintf.EscapeDefaultContext(settings.joinMethod), - "labels": strings.Join(labelsList, ","), - "databaseInstallMode": strconv.FormatBool(settings.databaseInstallMode), - "db_service_resource_labels": dbServiceResourceLabels, - "discoveryInstallMode": settings.discoveryInstallMode, - "discoveryGroup": shsprintf.EscapeDefaultContext(settings.discoveryGroup), - }) + installOpts, err := h.installScriptOptions(ctx) if err != nil { - return "", trace.Wrap(err) + return "", trace.Wrap(err, "Building install script options") } - return buf.String(), nil + nodeInstallOpts := scripts.InstallNodeScriptOptions{ + InstallOptions: installOpts, + Token: token.GetName(), + CAPins: caPins, + // We are using the joinMethod from the script settings instead of the one from the token + // to reproduce the previous script behavior. I'm also afraid that using the + // join method from the token would provide an oracle for an attacker wanting to discover + // the join method. + // We might want to change this in the future to lookup the join method from the token + // to avoid potential mismatch and allow the caller to not care about the join method. + JoinMethod: joinMethod, + Labels: token.GetSuggestedLabels(), + LabelMatchers: token.GetSuggestedAgentMatcherLabels(), + AppServiceEnabled: settings.appInstallMode, + AppName: settings.appName, + AppURI: settings.appURI, + DatabaseServiceEnabled: settings.databaseInstallMode, + DiscoveryServiceEnabled: settings.discoveryInstallMode, + DiscoveryGroup: settings.discoveryGroup, + } + + return scripts.GetNodeInstallScript(ctx, nodeInstallOpts) } // validateJoinToken validate a join token. @@ -789,17 +653,3 @@ func isSameAzureRuleSet(r1, r2 []*types.ProvisionTokenSpecV2Azure_Rule) bool { sortAzureRules(r2) return reflect.DeepEqual(r1, r2) } - -type nodeAPIGetter interface { - // GetToken looks up a provisioning token. - GetToken(ctx context.Context, token string) (types.ProvisionToken, error) - - // GetClusterCACert returns the CAs for the local cluster without signing keys. - GetClusterCACert(ctx context.Context) (*proto.GetClusterCACertResponse, error) - - // GetProxies returns a list of registered proxies. - GetProxies() ([]types.Server, error) -} - -// appURIPattern is a regexp excluding invalid characters from application URIs. -var appURIPattern = regexp.MustCompile(`^[-\w/:. ]+$`) diff --git a/lib/web/join_tokens_test.go b/lib/web/join_tokens_test.go index 4e0062b333e..515ea0d2113 100644 --- a/lib/web/join_tokens_test.go +++ b/lib/web/join_tokens_test.go @@ -23,27 +23,33 @@ import ( "encoding/hex" "encoding/json" "fmt" + "math/rand/v2" "net/http" "net/url" "regexp" + "strconv" "testing" "time" "github.com/google/go-cmp/cmp" "github.com/google/go-cmp/cmp/cmpopts" "github.com/gravitational/trace" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/client/proto" + autoupdatev1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/autoupdate/v1" "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils" + apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/lib/auth/authclient" + "github.com/gravitational/teleport/lib/automaticupgrades" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/modules" "github.com/gravitational/teleport/lib/services" libui "github.com/gravitational/teleport/lib/ui" + utils "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/web/ui" ) @@ -669,41 +675,18 @@ func toHex(s string) string { return hex.EncodeToString([]byte(s)) } func TestGetNodeJoinScript(t *testing.T) { validToken := "f18da1c9f6630a51e8daf121e7451daa" + invalidToken := "f18da1c9f6630a51e8daf121e7451dab" validIAMToken := "valid-iam-token" internalResourceID := "967d38ff-7a61-4f42-bd2d-c61965b44db0" - m := &mockedNodeAPIGetter{ - mockGetProxyServers: func() ([]types.Server, error) { - var s types.ServerV2 - s.SetPublicAddrs([]string{"test-host:12345678"}) - - return []types.Server{&s}, nil - }, - mockGetClusterCACert: func(context.Context) (*proto.GetClusterCACertResponse, error) { - fakeBytes := []byte(fixtures.SigningCertPEM) - return &proto.GetClusterCACertResponse{TLSCA: fakeBytes}, nil - }, - mockGetToken: func(_ context.Context, token string) (types.ProvisionToken, error) { - if token == validToken || token == validIAMToken { - return &types.ProvisionTokenV2{ - Metadata: types.Metadata{ - Name: token, - }, - Spec: types.ProvisionTokenSpecV2{ - SuggestedLabels: types.Labels{ - types.InternalResourceIDLabel: utils.Strings{internalResourceID}, - }, - }, - }, nil - } - return nil, trace.NotFound("token does not exist") - }, - } + hostname := "proxy.example.com" + port := 1234 for _, test := range []struct { desc string settings scriptSettings errAssert require.ErrorAssertionFunc + token *types.ProvisionTokenV2 extraAssertions func(script string) }{ { @@ -713,22 +696,52 @@ func TestGetNodeJoinScript(t *testing.T) { }, { desc: "short token length", - settings: scriptSettings{token: toHex("f18da1c9f6630a51e8daf121e7451d")}, + settings: scriptSettings{token: toHex(validToken[:30])}, errAssert: require.Error, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken[:30], + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, }, { desc: "valid length but does not exist", - settings: scriptSettings{token: toHex("xxxxxxx9f6630a51e8daf121exxxxxxx")}, + settings: scriptSettings{token: toHex(invalidToken)}, errAssert: require.Error, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, }, { desc: "valid", settings: scriptSettings{token: validToken}, errAssert: require.NoError, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, extraAssertions: func(script string) { require.Contains(t, script, validToken) - require.Contains(t, script, "test-host") - require.Contains(t, script, "12345678") + require.Contains(t, script, hostname) + require.Contains(t, script, strconv.Itoa(port)) require.Contains(t, script, "sha256:") require.NotContains(t, script, "JOIN_METHOD='iam'") }, @@ -747,6 +760,16 @@ func TestGetNodeJoinScript(t *testing.T) { token: validIAMToken, joinMethod: string(types.JoinMethodIAM), }, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validIAMToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, errAssert: require.NoError, extraAssertions: func(script string) { require.Contains(t, script, "JOIN_METHOD='iam'") @@ -756,14 +779,34 @@ func TestGetNodeJoinScript(t *testing.T) { desc: "internal resourceid label", settings: scriptSettings{token: validToken}, errAssert: require.NoError, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, extraAssertions: func(script string) { require.Contains(t, script, "--labels ") require.Contains(t, script, fmt.Sprintf("%s=%s", types.InternalResourceIDLabel, internalResourceID)) }, }, { - desc: "app server labels", - settings: scriptSettings{token: validToken, appInstallMode: true, appName: "app-name", appURI: "app-uri"}, + desc: "app server labels", + settings: scriptSettings{token: validToken, appInstallMode: true, appName: "app-name", appURI: "app-uri"}, + token: &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + }, + }, errAssert: require.NoError, extraAssertions: func(script string) { require.Contains(t, script, `APP_NAME='app-name'`) @@ -774,7 +817,12 @@ func TestGetNodeJoinScript(t *testing.T) { }, } { t.Run(test.desc, func(t *testing.T) { - script, err := getJoinScript(context.Background(), test.settings, m) + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + hostname: hostname, + port: port, + token: test.token, + }) + script, err := h.getJoinScript(context.Background(), test.settings) test.errAssert(t, err) if err != nil { require.Empty(t, script) @@ -787,28 +835,95 @@ func TestGetNodeJoinScript(t *testing.T) { } } +type autoupdateAccessPointMock struct { + authclient.ProxyAccessPoint + mock.Mock +} + +func (a *autoupdateAccessPointMock) GetAutoUpdateAgentRollout(ctx context.Context) (*autoupdatev1pb.AutoUpdateAgentRollout, error) { + args := a.Called(ctx) + return args.Get(0).(*autoupdatev1pb.AutoUpdateAgentRollout), args.Error(1) +} + +type autoupdateProxyClientMock struct { + authclient.ClientI + mock.Mock +} + +func (a *autoupdateProxyClientMock) GetToken(ctx context.Context, token string) (types.ProvisionToken, error) { + args := a.Called(ctx, token) + return args.Get(0).(types.ProvisionToken), args.Error(1) +} + +func (a *autoupdateProxyClientMock) GetClusterCACert(ctx context.Context) (*proto.GetClusterCACertResponse, error) { + args := a.Called(ctx) + return args.Get(0).(*proto.GetClusterCACertResponse), args.Error(1) +} + +type autoupdateTestHandlerConfig struct { + testModules *modules.TestModules + hostname string + port int + channels automaticupgrades.Channels + rollout *autoupdatev1pb.AutoUpdateAgentRollout + token *types.ProvisionTokenV2 +} + +func newAutoupdateTestHandler(t *testing.T, config autoupdateTestHandlerConfig) *Handler { + if config.hostname == "" { + config.hostname = fmt.Sprintf("proxy-%d.example.com", rand.Int()) + } + if config.port == 0 { + config.port = rand.IntN(65535) + } + addr := config.hostname + ":" + strconv.Itoa(config.port) + + if config.channels == nil { + config.channels = automaticupgrades.Channels{} + } + require.NoError(t, config.channels.CheckAndSetDefaults()) + + ap := &autoupdateAccessPointMock{} + if config.rollout == nil { + ap.On("GetAutoUpdateAgentRollout", mock.Anything).Return(config.rollout, trace.NotFound("rollout does not exist")) + } else { + ap.On("GetAutoUpdateAgentRollout", mock.Anything).Return(config.rollout, nil) + } + + clt := &autoupdateProxyClientMock{} + if config.token == nil { + clt.On("GetToken", mock.Anything, mock.Anything).Return(config.token, trace.NotFound("token does not exist")) + } else { + clt.On("GetToken", mock.Anything, config.token.GetName()).Return(config.token, nil) + } + + clt.On("GetClusterCACert", mock.Anything).Return(&proto.GetClusterCACertResponse{TLSCA: []byte(fixtures.SigningCertPEM)}, nil) + + if config.testModules == nil { + config.testModules = &modules.TestModules{ + TestBuildType: modules.BuildCommunity, + } + } + modules.SetTestModules(t, config.testModules) + h := &Handler{ + clusterFeatures: *config.testModules.Features().ToProto(), + cfg: Config{ + AutomaticUpgradesChannels: config.channels, + AccessPoint: ap, + PublicProxyAddr: addr, + ProxyClient: clt, + }, + logger: utils.NewSlogLoggerForTests(), + } + h.PublicProxyAddr() + return h +} + func TestGetAppJoinScript(t *testing.T) { testTokenID := "f18da1c9f6630a51e8daf121e7451daa" - m := &mockedNodeAPIGetter{ - mockGetToken: func(_ context.Context, token string) (types.ProvisionToken, error) { - if token == testTokenID { - return &types.ProvisionTokenV2{ - Metadata: types.Metadata{ - Name: token, - }, - }, nil - } - return nil, trace.NotFound("token does not exist") - }, - mockGetProxyServers: func() ([]types.Server, error) { - var s types.ServerV2 - s.SetPublicAddrs([]string{"test-host:12345678"}) - - return []types.Server{&s}, nil - }, - mockGetClusterCACert: func(context.Context) (*proto.GetClusterCACertResponse, error) { - fakeBytes := []byte(fixtures.SigningCertPEM) - return &proto.GetClusterCACertResponse{TLSCA: fakeBytes}, nil + token := &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: testTokenID, }, } badAppName := scriptSettings{ @@ -825,20 +940,24 @@ func TestGetAppJoinScript(t *testing.T) { appURI: "", } + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{token: token}) + hostname, port, err := utils.SplitHostPort(h.PublicProxyAddr()) + require.NoError(t, err) + // Test invalid app data. - script, err := getJoinScript(context.Background(), badAppName, m) + script, err := h.getJoinScript(context.Background(), badAppName) require.Empty(t, script) require.True(t, trace.IsBadParameter(err)) - script, err = getJoinScript(context.Background(), badAppURI, m) + script, err = h.getJoinScript(context.Background(), badAppURI) require.Empty(t, script) require.True(t, trace.IsBadParameter(err)) // Test various 'good' cases. expectedOutputs := []string{ testTokenID, - "test-host", - "12345678", + hostname, + port, "sha256:", } @@ -959,7 +1078,7 @@ func TestGetAppJoinScript(t *testing.T) { for _, tc := range tests { tc := tc t.Run(tc.desc, func(t *testing.T) { - script, err = getJoinScript(context.Background(), tc.settings, m) + script, err = h.getJoinScript(context.Background(), tc.settings) if tc.shouldError { require.Error(t, err) require.Empty(t, script) @@ -977,53 +1096,46 @@ func TestGetDatabaseJoinScript(t *testing.T) { validToken := "f18da1c9f6630a51e8daf121e7451daa" emptySuggestedAgentMatcherLabelsToken := "f18da1c9f6630a51e8daf121e7451000" internalResourceID := "967d38ff-7a61-4f42-bd2d-c61965b44db0" + hostname := "test.example.com" + port := 1234 - m := &mockedNodeAPIGetter{ - mockGetProxyServers: func() ([]types.Server, error) { - var s types.ServerV2 - s.SetPublicAddrs([]string{"test-host:12345678"}) + token := &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, + }, + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, + SuggestedAgentMatcherLabels: types.Labels{ + "env": apiutils.Strings{"prod"}, + "product": apiutils.Strings{"*"}, + "os": apiutils.Strings{"mac", "linux"}, + }, + }, + } - return []types.Server{&s}, nil + noMatcherToken := &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: emptySuggestedAgentMatcherLabelsToken, }, - mockGetClusterCACert: func(context.Context) (*proto.GetClusterCACertResponse, error) { - fakeBytes := []byte(fixtures.SigningCertPEM) - return &proto.GetClusterCACertResponse{TLSCA: fakeBytes}, nil - }, - mockGetToken: func(_ context.Context, token string) (types.ProvisionToken, error) { - provisionToken := &types.ProvisionTokenV2{ - Metadata: types.Metadata{ - Name: token, - }, - Spec: types.ProvisionTokenSpecV2{ - SuggestedLabels: types.Labels{ - types.InternalResourceIDLabel: utils.Strings{internalResourceID}, - }, - SuggestedAgentMatcherLabels: types.Labels{ - "env": utils.Strings{"prod"}, - "product": utils.Strings{"*"}, - "os": utils.Strings{"mac", "linux"}, - }, - }, - } - if token == validToken { - return provisionToken, nil - } - if token == emptySuggestedAgentMatcherLabelsToken { - provisionToken.Spec.SuggestedAgentMatcherLabels = types.Labels{} - return provisionToken, nil - } - return nil, trace.NotFound("token does not exist") + Spec: types.ProvisionTokenSpecV2{ + SuggestedLabels: types.Labels{ + types.InternalResourceIDLabel: apiutils.Strings{internalResourceID}, + }, }, } for _, test := range []struct { desc string settings scriptSettings + token *types.ProvisionTokenV2 errAssert require.ErrorAssertionFunc extraAssertions func(script string) }{ { - desc: "two installation methods", + desc: "two installation methods", + token: token, settings: scriptSettings{ token: validToken, databaseInstallMode: true, @@ -1032,7 +1144,8 @@ func TestGetDatabaseJoinScript(t *testing.T) { errAssert: require.Error, }, { - desc: "valid", + desc: "valid", + token: token, settings: scriptSettings{ databaseInstallMode: true, token: validToken, @@ -1040,7 +1153,8 @@ func TestGetDatabaseJoinScript(t *testing.T) { errAssert: require.NoError, extraAssertions: func(script string) { require.Contains(t, script, validToken) - require.Contains(t, script, "test-host") + require.Contains(t, script, hostname) + require.Contains(t, script, strconv.Itoa(port)) require.Contains(t, script, "sha256:") require.Contains(t, script, "--labels ") require.Contains(t, script, fmt.Sprintf("%s=%s", types.InternalResourceIDLabel, internalResourceID)) @@ -1058,7 +1172,8 @@ db_service: }, }, { - desc: "empty suggestedAgentMatcherLabels", + desc: "empty suggestedAgentMatcherLabels", + token: noMatcherToken, settings: scriptSettings{ databaseInstallMode: true, token: emptySuggestedAgentMatcherLabelsToken, @@ -1066,7 +1181,8 @@ db_service: errAssert: require.NoError, extraAssertions: func(script string) { require.Contains(t, script, emptySuggestedAgentMatcherLabelsToken) - require.Contains(t, script, "test-host") + require.Contains(t, script, hostname) + require.Contains(t, script, strconv.Itoa(port)) require.Contains(t, script, "sha256:") require.Contains(t, script, "--labels ") require.Contains(t, script, fmt.Sprintf("%s=%s", types.InternalResourceIDLabel, internalResourceID)) @@ -1081,7 +1197,13 @@ db_service: }, } { t.Run(test.desc, func(t *testing.T) { - script, err := getJoinScript(context.Background(), test.settings, m) + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + hostname: hostname, + port: port, + token: test.token, + }) + + script, err := h.getJoinScript(context.Background(), test.settings) test.errAssert(t, err) if err != nil { require.Empty(t, script) @@ -1096,30 +1218,13 @@ db_service: func TestGetDiscoveryJoinScript(t *testing.T) { const validToken = "f18da1c9f6630a51e8daf121e7451daa" - - m := &mockedNodeAPIGetter{ - mockGetProxyServers: func() ([]types.Server, error) { - var s types.ServerV2 - s.SetPublicAddrs([]string{"test-host:12345678"}) - - return []types.Server{&s}, nil - }, - mockGetClusterCACert: func(context.Context) (*proto.GetClusterCACertResponse, error) { - fakeBytes := []byte(fixtures.SigningCertPEM) - return &proto.GetClusterCACertResponse{TLSCA: fakeBytes}, nil - }, - mockGetToken: func(_ context.Context, token string) (types.ProvisionToken, error) { - provisionToken := &types.ProvisionTokenV2{ - Metadata: types.Metadata{ - Name: token, - }, - Spec: types.ProvisionTokenSpecV2{}, - } - if token == validToken { - return provisionToken, nil - } - return nil, trace.NotFound("token does not exist") + hostname := "test.example.com" + port := 1234 + token := &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, }, + Spec: types.ProvisionTokenSpecV2{}, } for _, test := range []struct { @@ -1138,7 +1243,8 @@ func TestGetDiscoveryJoinScript(t *testing.T) { errAssert: require.NoError, extraAssertions: func(t *testing.T, script string) { require.Contains(t, script, validToken) - require.Contains(t, script, "test-host") + require.Contains(t, script, hostname) + require.Contains(t, script, strconv.Itoa(port)) require.Contains(t, script, "sha256:") require.Contains(t, script, "--labels ") require.Contains(t, script, ` @@ -1157,7 +1263,12 @@ discovery_service: }, } { t.Run(test.desc, func(t *testing.T) { - script, err := getJoinScript(context.Background(), test.settings, m) + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + hostname: hostname, + port: port, + token: token, + }) + script, err := h.getJoinScript(context.Background(), test.settings) test.errAssert(t, err) if err != nil { require.Empty(t, script) @@ -1276,28 +1387,9 @@ func TestIsSameRuleSet(t *testing.T) { func TestJoinScript(t *testing.T) { validToken := "f18da1c9f6630a51e8daf121e7451daa" - - m := &mockedNodeAPIGetter{ - mockGetProxyServers: func() ([]types.Server, error) { - return []types.Server{ - &types.ServerV2{ - Spec: types.ServerSpecV2{ - PublicAddrs: []string{"test-host:12345678"}, - Version: teleport.Version, - }, - }, - }, nil - }, - mockGetClusterCACert: func(context.Context) (*proto.GetClusterCACertResponse, error) { - fakeBytes := []byte(fixtures.SigningCertPEM) - return &proto.GetClusterCACertResponse{TLSCA: fakeBytes}, nil - }, - mockGetToken: func(_ context.Context, token string) (types.ProvisionToken, error) { - return &types.ProvisionTokenV2{ - Metadata: types.Metadata{ - Name: token, - }, - }, nil + token := &types.ProvisionTokenV2{ + Metadata: types.Metadata{ + Name: validToken, }, } @@ -1305,8 +1397,11 @@ func TestJoinScript(t *testing.T) { getGravitationalTeleportLinkRegex := regexp.MustCompile(`https://cdn\.teleport\.dev/\${TELEPORT_PACKAGE_NAME}[-_]v?\${TELEPORT_VERSION}`) t.Run("oss", func(t *testing.T) { + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + token: token, + }) // Using the OSS Version, all the links must contain only teleport as package name. - script, err := getJoinScript(context.Background(), scriptSettings{token: validToken}, m) + script, err := h.getJoinScript(context.Background(), scriptSettings{token: validToken}) require.NoError(t, err) matches := getGravitationalTeleportLinkRegex.FindAllString(script, -1) @@ -1321,8 +1416,11 @@ func TestJoinScript(t *testing.T) { t.Run("ent", func(t *testing.T) { // Using the Enterprise Version, the package name must be teleport-ent - modules.SetTestModules(t, &modules.TestModules{TestBuildType: modules.BuildEnterprise}) - script, err := getJoinScript(context.Background(), scriptSettings{token: validToken}, m) + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + testModules: &modules.TestModules{TestBuildType: modules.BuildEnterprise}, + token: token, + }) + script, err := h.getJoinScript(context.Background(), scriptSettings{token: validToken}) require.NoError(t, err) matches := getGravitationalTeleportLinkRegex.FindAllString(script, -1) @@ -1338,8 +1436,16 @@ func TestJoinScript(t *testing.T) { t.Run("using repo", func(t *testing.T) { t.Run("installUpdater is true", func(t *testing.T) { - currentStableCloudVersion := "v99.1.1" - script, err := getJoinScript(context.Background(), scriptSettings{token: validToken, installUpdater: true, automaticUpgradesVersion: currentStableCloudVersion}, m) + currentStableCloudVersion := "1.2.3" + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + testModules: &modules.TestModules{TestFeatures: modules.Features{Cloud: true, AutomaticUpgrades: true}}, + token: token, + channels: automaticupgrades.Channels{ + automaticupgrades.DefaultChannelName: &automaticupgrades.Channel{StaticVersion: currentStableCloudVersion}, + }, + }) + + script, err := h.getJoinScript(context.Background(), scriptSettings{token: validToken}) require.NoError(t, err) // list of packages must include the updater @@ -1356,10 +1462,13 @@ func TestJoinScript(t *testing.T) { // Repo channel is stable/cloud require.Contains(t, script, "REPO_CHANNEL='stable/cloud'") // TELEPORT_VERSION is the one provided by https://updates.releases.teleport.dev/v1/stable/cloud/version - require.Contains(t, script, "TELEPORT_VERSION='99.1.1'") + require.Contains(t, script, fmt.Sprintf("TELEPORT_VERSION='%s'", currentStableCloudVersion)) }) t.Run("installUpdater is false", func(t *testing.T) { - script, err := getJoinScript(context.Background(), scriptSettings{token: validToken, installUpdater: false}, m) + h := newAutoupdateTestHandler(t, autoupdateTestHandlerConfig{ + token: token, + }) + script, err := h.getJoinScript(context.Background(), scriptSettings{token: validToken}) require.NoError(t, err) require.Contains(t, script, ""+ " PACKAGE_LIST=${TELEPORT_PACKAGE_PIN_VERSION}\n"+ @@ -1484,32 +1593,3 @@ func TestIsSameAzureRuleSet(t *testing.T) { }) } } - -type mockedNodeAPIGetter struct { - mockGetProxyServers func() ([]types.Server, error) - mockGetClusterCACert func(ctx context.Context) (*proto.GetClusterCACertResponse, error) - mockGetToken func(ctx context.Context, token string) (types.ProvisionToken, error) -} - -func (m *mockedNodeAPIGetter) GetProxies() ([]types.Server, error) { - if m.mockGetProxyServers != nil { - return m.mockGetProxyServers() - } - - return nil, trace.NotImplemented("mockGetProxyServers not implemented") -} - -func (m *mockedNodeAPIGetter) GetClusterCACert(ctx context.Context) (*proto.GetClusterCACertResponse, error) { - if m.mockGetClusterCACert != nil { - return m.mockGetClusterCACert(ctx) - } - - return nil, trace.NotImplemented("mockGetClusterCACert not implemented") -} - -func (m *mockedNodeAPIGetter) GetToken(ctx context.Context, token string) (types.ProvisionToken, error) { - if m.mockGetToken != nil { - return m.mockGetToken(ctx, token) - } - return nil, trace.NotImplemented("mockGetToken not implemented") -} diff --git a/lib/web/scripts/install_node.go b/lib/web/scripts/install_node.go index 87fffd7b587..4c3cae2648c 100644 --- a/lib/web/scripts/install_node.go +++ b/lib/web/scripts/install_node.go @@ -19,19 +19,30 @@ package scripts import ( + "bytes" + "context" _ "embed" "fmt" + regexp "regexp" "sort" + "strconv" "strings" "text/template" + "github.com/google/safetext/shsprintf" "github.com/gravitational/trace" "gopkg.in/yaml.v3" + "k8s.io/apimachinery/pkg/util/validation" "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/utils" + apiutils "github.com/gravitational/teleport/api/utils" + "github.com/gravitational/teleport/lib/automaticupgrades" + "github.com/gravitational/teleport/lib/utils" ) +// appURIPattern is a regexp excluding invalid characters from application URIs. +var appURIPattern = regexp.MustCompile(`^[-\w/:. ]+$`) + // ErrorBashScript is used to display friendly error message when // there is an error prepping the actual script. var ErrorBashScript = []byte(` @@ -44,11 +55,152 @@ exit 1 // to install teleport and join a teleport cluster. // //go:embed node-join/install.sh -var installNodeBashScript string +var installNodeBashScriptRaw string -var InstallNodeBashScript = template.Must(template.New("nodejoin").Parse(installNodeBashScript)) +var installNodeBashScript = template.Must(template.New("nodejoin").Parse(installNodeBashScriptRaw)) -// MarshalLabelsYAML returns a list of strings, each one containing a +// InstallNodeScriptOptions contains the options configuring the install-node script. +type InstallNodeScriptOptions struct { + // Required for installation + InstallOptions InstallScriptOptions + + // Required for joining + Token string + CAPins []string + JoinMethod types.JoinMethod + + // Required for service configuration + Labels types.Labels + LabelMatchers types.Labels + + AppServiceEnabled bool + AppName string + AppURI string + + DatabaseServiceEnabled bool + DiscoveryServiceEnabled bool + DiscoveryGroup string +} + +// GetNodeInstallScript generates an agent installation script which will: +// - install Teleport +// - configure the Teleport agent joining +// - configure the Teleport agent services (currently support ssh, app, database, and discovery) +// - start the agent +func GetNodeInstallScript(ctx context.Context, opts InstallNodeScriptOptions) (string, error) { + // Computing installation-related values + + // By default, it will use `stable/v`, eg stable/v12 + repoChannel := "" + installPackageUpdater := false + + switch opts.InstallOptions.AutoupdateStyle { + case NoAutoupdate: + case PackageManagerAutoupdate: + // Note: This is a cloud-specific repo. We could use the new stable/rolling + // repo in non-cloud case, but the script has never support enabling autoupdates + // in a non-cloud cluster. + // We will prefer using the new updater binary for autoupdates in self-hosted setups. + repoChannel = automaticupgrades.DefaultCloudChannelName + installPackageUpdater = true + case UpdaterBinaryAutoupdate: + // TODO(hugoShaka): add new autoupdate binary support in the node-install script + // by using oneoff in another PR + return "", trace.NotImplemented("This path is not implemented yet.") + default: + return "", trace.BadParameter("unsupported autoupdate style: %v", opts.InstallOptions.AutoupdateStyle) + } + + // Computing joining-related values + hostname, portStr, err := utils.SplitHostPort(opts.InstallOptions.ProxyAddr) + if err != nil { + return "", trace.Wrap(err) + } + + // Computing service configuration-related values + labelsList := []string{} + for labelKey, labelValues := range opts.Labels { + labels := strings.Join(labelValues, " ") + labelsList = append(labelsList, fmt.Sprintf("%s=%s", labelKey, labels)) + } + + var dbServiceResourceLabels []string + if opts.DatabaseServiceEnabled { + dbServiceResourceLabels, err = marshalLabelsYAML(opts.LabelMatchers, 6) + if err != nil { + return "", trace.Wrap(err) + } + } + + var appServerResourceLabels []string + + if opts.AppServiceEnabled { + if errs := validation.IsDNS1035Label(opts.AppName); len(errs) > 0 { + return "", trace.BadParameter("appName %q must be a valid DNS subdomain: https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#application-name", opts.AppName) + } + if !appURIPattern.MatchString(opts.AppURI) { + return "", trace.BadParameter("appURI %q contains invalid characters", opts.AppURI) + } + + appServerResourceLabels, err = marshalLabelsYAML(opts.Labels, 4) + if err != nil { + return "", trace.Wrap(err) + } + } + + if opts.DiscoveryServiceEnabled { + if opts.DiscoveryGroup == "" { + return "", trace.BadParameter("discovery group is required") + } + } + + var buf bytes.Buffer + + // TODO(hugoShaka): burn this map and replace it by something saner in a future PR. + + // This section relies on Go's default zero values to make sure that the settings + // are correct when not installing an app. + err = installNodeBashScript.Execute(&buf, map[string]interface{}{ + "token": opts.Token, + "hostname": hostname, + "port": portStr, + // The install.sh script has some manually generated configs and some + // generated by the `teleport config` commands. The old bash + // version used space delimited values whereas the teleport command uses + // a comma delimeter. The Old version can be removed when the install.sh + // file has been completely converted over. + "caPinsOld": strings.Join(opts.CAPins, " "), + "caPins": strings.Join(opts.CAPins, ","), + "packageName": opts.InstallOptions.TeleportFlavor, + "repoChannel": repoChannel, + "installUpdater": strconv.FormatBool(installPackageUpdater), + "version": shsprintf.EscapeDefaultContext(opts.InstallOptions.TeleportVersion), + "appInstallMode": strconv.FormatBool(opts.AppServiceEnabled), + "appServerResourceLabels": appServerResourceLabels, + "appName": shsprintf.EscapeDefaultContext(opts.AppName), + "appURI": shsprintf.EscapeDefaultContext(opts.AppURI), + "joinMethod": shsprintf.EscapeDefaultContext(string(opts.JoinMethod)), + "labels": strings.Join(labelsList, ","), + "databaseInstallMode": strconv.FormatBool(opts.DatabaseServiceEnabled), + // No one knows why this field is in snake case ¯\_(ツ)_/¯ + // Also, even if the name is similar to appServerResourceLabels, they must not be confused. + // appServerResourceLabels are labels to apply on the declared app, while + // db_service_resource_labels are labels matchers for the service to select resources to serve. + "db_service_resource_labels": dbServiceResourceLabels, + "discoveryInstallMode": strconv.FormatBool(opts.DiscoveryServiceEnabled), + "discoveryGroup": shsprintf.EscapeDefaultContext(opts.DiscoveryGroup), + }) + if err != nil { + return "", trace.Wrap(err) + } + + return buf.String(), nil +} + +// TODO(hugoShaka): burn the indentation thing, this is too fragile and show be handled +// by the template itself. + +// marshalLabelsYAML returns a list of strings, each one containing a // label key and list of value's pair. // This is used to create yaml sections within the join scripts. // @@ -56,7 +208,7 @@ var InstallNodeBashScript = template.Must(template.New("nodejoin").Parse(install // top of the default space already used, for the default yaml listing // format (the listing values with the dashes). If `extraListIndent` // is zero, it's equivalent to using default space only (which is 4 spaces). -func MarshalLabelsYAML(resourceMatcherLabels types.Labels, extraListIndent int) ([]string, error) { +func marshalLabelsYAML(resourceMatcherLabels types.Labels, extraListIndent int) ([]string, error) { if len(resourceMatcherLabels) == 0 { return []string{"{}"}, nil } @@ -73,7 +225,7 @@ func MarshalLabelsYAML(resourceMatcherLabels types.Labels, extraListIndent int) for _, labelName := range labelKeys { labelValues := resourceMatcherLabels[labelName] - bs, err := yaml.Marshal(map[string]utils.Strings{labelName: labelValues}) + bs, err := yaml.Marshal(map[string]apiutils.Strings{labelName: labelValues}) if err != nil { return nil, trace.Wrap(err) } diff --git a/lib/web/scripts/install_node_test.go b/lib/web/scripts/install_node_test.go index f56a44546e7..141133c5b9b 100644 --- a/lib/web/scripts/install_node_test.go +++ b/lib/web/scripts/install_node_test.go @@ -66,7 +66,7 @@ func TestMarshalLabelsYAML(t *testing.T) { numExtraIndent: 2, }, } { - got, err := MarshalLabelsYAML(tt.labels, tt.numExtraIndent) + got, err := marshalLabelsYAML(tt.labels, tt.numExtraIndent) require.NoError(t, err) require.YAMLEq(t, strings.Join(tt.expected, "\n"), strings.Join(got, "\n"))