Skip to content

Commit b91442f

Browse files
fix(ske): address stateless login review feedback
1 parent b188de2 commit b91442f

4 files changed

Lines changed: 72 additions & 7 deletions

File tree

internal/cmd/ske/kubeconfig/login/login.go

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -109,7 +109,11 @@ func NewCmd(params *types.CmdParams) *cobra.Command {
109109
if err != nil {
110110
return err
111111
}
112-
tokenEndpoint, err := auth.GetIDPTokenEndpoint(params.Printer)
112+
getTokenEndpoint := auth.GetIDPTokenEndpoint
113+
if statelessIDPMode {
114+
getTokenEndpoint = auth.GetIDPTokenEndpointStateless
115+
}
116+
tokenEndpoint, err := getTokenEndpoint(params.Printer)
113117
if err != nil {
114118
return fmt.Errorf("get IDP token endpoint: %w", err)
115119
}
@@ -186,11 +190,11 @@ func parseClusterConfig(p *print.Printer, cmd *cobra.Command, idpMode, workloadI
186190
}
187191
}
188192

189-
if clusterName := flags.FlagToStringValue(p, cmd, clusterNameFlag); clusterName != "" {
190-
clusterConfig.ClusterName = clusterName
193+
if clusterConfig.ClusterName == "" {
194+
clusterConfig.ClusterName = flags.FlagToStringValue(p, cmd, clusterNameFlag)
191195
}
192-
if organizationID := flags.FlagToStringValue(p, cmd, organizationFlag); organizationID != "" {
193-
clusterConfig.OrganizationID = organizationID
196+
if clusterConfig.OrganizationID == "" {
197+
clusterConfig.OrganizationID = flags.FlagToStringValue(p, cmd, organizationFlag)
194198
}
195199
globalFlags := globalflags.Parse(p, cmd)
196200
if clusterConfig.STACKITProjectID == "" {

internal/cmd/ske/kubeconfig/login/login_test.go

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -115,6 +115,8 @@ func TestParseClusterConfigWithoutExecClusterInfo(t *testing.T) {
115115
func TestParseClusterConfigUsesExecClusterInfo(t *testing.T) {
116116
viper.Reset()
117117
t.Cleanup(viper.Reset)
118+
viper.Set(config.ProjectIdKey, uuid.NewString())
119+
viper.Set(config.RegionKey, "conflicting-region")
118120
t.Setenv(envServiceAccountEmail, "workload@sa.stackit.cloud")
119121

120122
configJSON, err := json.Marshal(fixtureClusterConfig())
@@ -129,6 +131,12 @@ func TestParseClusterConfigUsesExecClusterInfo(t *testing.T) {
129131
params := testparams.NewTestParams()
130132
cmd := &cobra.Command{}
131133
configureFlags(cmd)
134+
if err := cmd.Flags().Set(clusterNameFlag, "conflicting-cluster"); err != nil {
135+
t.Fatalf("set cluster name flag: %v", err)
136+
}
137+
if err := cmd.Flags().Set(organizationFlag, uuid.NewString()); err != nil {
138+
t.Fatalf("set organization flag: %v", err)
139+
}
132140
actual, err := parseClusterConfig(params.Printer, cmd, true, true, true)
133141
if err != nil {
134142
t.Fatalf("parse cluster config: %v", err)

internal/pkg/auth/utils.go

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -71,14 +71,20 @@ func retrieveIDPWellKnownConfigWithStorage(p *print.Printer, persistTokenEndpoin
7171
return idpWellKnownConfig, nil
7272
}
7373

74-
// GetIDPTokenEndpoint returns the configured IdP token endpoint without requiring
75-
// authentication storage. Persisted configuration is reused when available.
74+
// GetIDPTokenEndpoint returns the persisted IdP token endpoint when available and
75+
// falls back to endpoint discovery.
7676
func GetIDPTokenEndpoint(p *print.Printer) (string, error) {
7777
tokenEndpoint, err := GetAuthField(IDP_TOKEN_ENDPOINT)
7878
if err == nil && tokenEndpoint != "" {
7979
return tokenEndpoint, nil
8080
}
8181

82+
return GetIDPTokenEndpointStateless(p)
83+
}
84+
85+
// GetIDPTokenEndpointStateless discovers the IdP token endpoint without reading
86+
// from or writing to authentication storage.
87+
func GetIDPTokenEndpointStateless(p *print.Printer) (string, error) {
8288
wellKnownConfig, err := retrieveIDPWellKnownConfigWithStorage(p, false)
8389
if err != nil {
8490
return "", err

internal/pkg/auth/utils_test.go

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,16 @@
11
package auth
22

33
import (
4+
"net/http"
5+
"net/http/httptest"
46
"testing"
57

68
"github.com/google/go-cmp/cmp"
79
"github.com/spf13/viper"
810
"github.com/zalando/go-keyring"
911

1012
"github.com/stackitcloud/stackit-cli/internal/pkg/config"
13+
"github.com/stackitcloud/stackit-cli/internal/pkg/testparams"
1114
)
1215

1316
func TestGetWellKnownConfig(t *testing.T) {
@@ -229,3 +232,47 @@ func TestParseWellKnownConfigWithoutStorage(t *testing.T) {
229232
t.Fatalf("Expected token endpoint not to be changed, got %q", storedTokenEndpoint)
230233
}
231234
}
235+
236+
func TestGetIDPTokenEndpointStatelessIgnoresStorage(t *testing.T) {
237+
viper.Reset()
238+
t.Cleanup(viper.Reset)
239+
keyring.MockInit()
240+
241+
const storedEndpoint = "https://stored.stackit.cloud/oauth"
242+
if err := SetAuthField(IDP_TOKEN_ENDPOINT, storedEndpoint); err != nil {
243+
t.Fatalf("Set stored token endpoint: %v", err)
244+
}
245+
246+
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
247+
w.Header().Set("Content-Type", "application/json")
248+
_, _ = w.Write([]byte(`{"issuer":"https://issuer.stackit.cloud/endpoint","authorization_endpoint":"https://auth.stackit.cloud/endpoint","token_endpoint":"https://discovered.stackit.cloud/oauth"}`))
249+
}))
250+
defer server.Close()
251+
252+
originalTransport := http.DefaultTransport
253+
http.DefaultTransport = server.Client().Transport
254+
t.Cleanup(func() {
255+
http.DefaultTransport = originalTransport
256+
})
257+
viper.Set(config.IdentityProviderCustomWellKnownConfigurationKey, server.URL)
258+
viper.Set(config.AllowedUrlDomainKey, "")
259+
260+
params := testparams.NewTestParams()
261+
params.Printer.AssumeYes = true
262+
actual, err := GetIDPTokenEndpointStateless(params.Printer)
263+
if err != nil {
264+
t.Fatalf("Get stateless token endpoint: %v", err)
265+
}
266+
const discoveredEndpoint = "https://discovered.stackit.cloud/oauth"
267+
if actual != discoveredEndpoint {
268+
t.Fatalf("Expected discovered endpoint %q, got %q", discoveredEndpoint, actual)
269+
}
270+
271+
actualStoredEndpoint, err := GetAuthField(IDP_TOKEN_ENDPOINT)
272+
if err != nil {
273+
t.Fatalf("Get stored token endpoint: %v", err)
274+
}
275+
if actualStoredEndpoint != storedEndpoint {
276+
t.Fatalf("Expected stored endpoint to remain %q, got %q", storedEndpoint, actualStoredEndpoint)
277+
}
278+
}

0 commit comments

Comments
 (0)