Skip to content
Merged
43 changes: 26 additions & 17 deletions cmd/kubectl.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,7 @@ import (
"os/exec"
"strings"

"github.com/rancher/norman/clientbase"
client "github.com/rancher/rancher/pkg/client/generated/management/v3"
extv1 "github.com/rancher/rancher/pkg/apis/ext.cattle.io/v1"
"github.com/urfave/cli/v3"
"k8s.io/client-go/tools/clientcmd"
"k8s.io/client-go/tools/clientcmd/api"
Expand Down Expand Up @@ -53,12 +52,33 @@ func runKubectl(ctx context.Context, cmd *cli.Command) error {
}

currentToken := currentRancherServer.AccessKey
t, err := c.ManagementClient.Token.ByID(currentToken)
// bearerToken intentionally keeps any "ext/" prefix on AccessKey because
// the ext API authenticator parses the prefix back out of the Bearer header.
bearerToken := currentRancherServer.AccessKey + ":" + currentRancherServer.SecretKey

// Norman's MasterClient does not expose its underlying http.Client, so
// build a parallel one here for the direct ext API call.
tlsConf, err := getTLSConfig(false, currentRancherServer.CACerts)
if err != nil {
return err
return fmt.Errorf("error creating TLS config: %w", err)
}
httpClient, err := newHTTPClient(currentRancherServer, tlsConf)
if err != nil {
return fmt.Errorf("error creating HTTP client: %w", err)
}
baseURL, err := currentRancherServer.EnvironmentURL()
if err != nil {
return fmt.Errorf("error resolving server base URL: %w", err)
}
extGetter := func(ctx context.Context, id string) (*extv1.Token, error) {
return getExtToken(ctx, id, baseURL, bearerToken, httpClient)
}
v3ByID := c.ManagementClient.Token.ByID

currentUser := t.UserID
currentUser, err := getTokenUserID(ctx, currentToken, v3ByID, extGetter)
if err != nil {
return err
}
kubeConfig, err := getKubeConfigForUser(cmd, currentUser)
if err != nil {
return err
Expand All @@ -70,7 +90,7 @@ func runKubectl(ctx context.Context, cmd *cli.Command) error {
if err != nil {
return err
}
isTokenValid, err = validateToken(tokenID, c.ManagementClient.Token)
isTokenValid, err = validateToken(ctx, tokenID, v3ByID, extGetter)
if err != nil {
return err
}
Expand Down Expand Up @@ -137,14 +157,3 @@ func extractKubeconfigTokenID(kubeconfig api.Config) (string, error) {

return parts[0], nil
}

func validateToken(tokenID string, tokenClient client.TokenOperations) (bool, error) {
token, err := tokenClient.ByID(tokenID)
if err != nil {
if !clientbase.IsNotFound(err) {
return false, err
}
return false, nil
}
return !token.Expired, nil
}
40 changes: 22 additions & 18 deletions cmd/kubectl_token.go
Original file line number Diff line number Diff line change
Expand Up @@ -339,24 +339,31 @@ func loginAndGenerateCred(client *http.Client, input *LoginInput) (*config.ExecC
}
input.authProvider = selectedProvider.GetType()

token := managementClient.Token{}
if samlProviders[input.authProvider] {
token, err = samlAuth(client, input, useV1Public)
var token loginToken
switch {
case samlProviders[input.authProvider]:
samlTok, err := samlAuth(client, input, useV1Public)
if err != nil {
return nil, err
}
} else if oauthProviders[input.authProvider] {
token = loginToken{
BearerToken: samlTok.Token,
ExpiresAt: samlTok.ExpiresAt,
UserID: samlTok.UserID,
}
case oauthProviders[input.authProvider]:
tokenPtr, err := oauthAuth(client, input, selectedProvider, useV1Public)
if err != nil {
return nil, err
}
token = *tokenPtr
} else {
default:
customPrint(fmt.Sprintf("Enter credentials for %s \n", input.authProvider))
token, err = basicAuth(client, input, useV1Public)
tok, err := basicAuth(client, input, useV1Public)
if err != nil {
return nil, err
}
token = tok
}

cred := &config.ExecCredential{
Expand All @@ -366,7 +373,7 @@ func loginAndGenerateCred(client *http.Client, input *LoginInput) (*config.ExecC
},
Status: &config.ExecCredentialStatus{},
}
cred.Status.Token = token.Token
cred.Status.Token = token.BearerToken
if token.ExpiresAt == "" {
return cred, nil
}
Expand All @@ -377,19 +384,16 @@ func loginAndGenerateCred(client *http.Client, input *LoginInput) (*config.ExecC
}
cred.Status.ExpirationTimestamp = &config.Time{Time: ts}
return cred, nil

}

func basicAuth(client *http.Client, input *LoginInput, useV1Public bool) (managementClient.Token, error) {
token := managementClient.Token{}

func basicAuth(client *http.Client, input *LoginInput, useV1Public bool) (loginToken, error) {
prompt := "Enter username"
if input.userID != "" {
prompt += " [" + input.userID + "]"
}
username, err := customPrompt(prompt+": ", true)
if err != nil {
return token, err
return loginToken{}, err
}

if username == "" && input.userID != "" {
Expand All @@ -398,7 +402,7 @@ func basicAuth(client *http.Client, input *LoginInput, useV1Public bool) (manage

password, err := customPrompt("Enter password: ", false)
if err != nil {
return token, err
return loginToken{}, err
}

responseType := "kubeconfig"
Expand All @@ -413,7 +417,7 @@ func basicAuth(client *http.Client, input *LoginInput, useV1Public bool) (manage
"password": password,
})
if err != nil {
return token, fmt.Errorf("failed to marshal request body: %w", err)
return loginToken{}, fmt.Errorf("failed to marshal request body: %w", err)
}

reqURL := fmt.Sprintf(loginURL, input.server)
Expand All @@ -424,20 +428,20 @@ func basicAuth(client *http.Client, input *LoginInput, useV1Public bool) (manage

req, err := http.NewRequest(http.MethodPost, reqURL, bytes.NewReader(reqBody))
if err != nil {
return token, fmt.Errorf("error creating request: %w", err)
return loginToken{}, fmt.Errorf("error creating request: %w", err)
}

resp, respBody, err := doRequest(client, req)
if err == nil && resp.StatusCode != http.StatusCreated {
err = fmt.Errorf("%d %s", resp.StatusCode, http.StatusText(resp.StatusCode))
}
if err != nil {
return token, fmt.Errorf("error logging user in: %w", err)
return loginToken{}, fmt.Errorf("error logging user in: %w", err)
}

err = json.Unmarshal(respBody, &token)
token, err := parseLoginResponse(respBody)
if err != nil {
return token, fmt.Errorf("error unmarshaling login response: %w", err)
return loginToken{}, err
}

return token, nil
Expand Down
16 changes: 7 additions & 9 deletions cmd/kubectl_token_oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@ import (
"time"

apiv3 "github.com/rancher/rancher/pkg/apis/management.cattle.io/v3"
managementClient "github.com/rancher/rancher/pkg/client/generated/management/v3"
"github.com/sirupsen/logrus"
"golang.org/x/oauth2"
)
Expand All @@ -33,7 +32,7 @@ const (
)

// oauthAuth dispatches the OAuth authentication flow based on the auth flow type.
func oauthAuth(client *http.Client, input *LoginInput, provider TypedProvider, useV1Public bool) (*managementClient.Token, error) {
func oauthAuth(client *http.Client, input *LoginInput, provider TypedProvider, useV1Public bool) (*loginToken, error) {
if input.authFlow == "" { // The flag has precedence over the env variable.
input.authFlow = os.Getenv("CATTLE_OAUTH_AUTH_FLOW")
}
Expand Down Expand Up @@ -77,7 +76,7 @@ func oauthAuthCodeAuth(
timeoutAfter time.Duration,
useV1Public bool,
openBrowser openBrowserFunc,
) (*managementClient.Token, error) {
) (*loginToken, error) {
oauthConfig, err := newOauthConfig(provider)
if err != nil {
return nil, fmt.Errorf("failed to create oauth config: %w", err)
Expand Down Expand Up @@ -271,7 +270,7 @@ func startCallbackServer(listener net.Listener, expectedState string, resultCh c
}

// oauthDeviceCodeAuth implements the device code flow for OAuth authentication.
func oauthDeviceCodeAuth(client *http.Client, input *LoginInput, provider TypedProvider, useV1Public bool) (*managementClient.Token, error) {
func oauthDeviceCodeAuth(client *http.Client, input *LoginInput, provider TypedProvider, useV1Public bool) (*loginToken, error) {
oauthConfig, err := newOauthConfig(provider)
if err != nil {
return nil, fmt.Errorf("failed to create oauth config: %w", err)
Expand Down Expand Up @@ -325,7 +324,7 @@ func newOauthConfig(provider TypedProvider) (*oauth2.Config, error) {
}

// rancherLogin sends the obtained OAuth token to Rancher to exchange it for a Rancher token that can be used for API authentication.
func rancherLogin(client *http.Client, input *LoginInput, oauthToken *oauth2.Token, useV1Public bool) (*managementClient.Token, error) {
func rancherLogin(client *http.Client, input *LoginInput, oauthToken *oauth2.Token, useV1Public bool) (*loginToken, error) {
reqURL := fmt.Sprintf(loginURL, input.server)
if !useV1Public {
providerName := strings.ToLower(strings.TrimSuffix(input.authProvider, "Provider"))
Expand Down Expand Up @@ -360,11 +359,10 @@ func rancherLogin(client *http.Client, input *LoginInput, oauthToken *oauth2.Tok
return nil, err
}

token := &managementClient.Token{}
err = json.Unmarshal(respBody, token)
token, err := parseLoginResponse(respBody)
if err != nil {
return nil, fmt.Errorf("error unmarshaling login response: %w", err)
return nil, err
}

return token, nil
return &token, nil
}
40 changes: 32 additions & 8 deletions cmd/kubectl_token_oauth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -153,12 +153,13 @@ func TestRancherLogin(t *testing.T) {
expiresAt := time.Now().Add(time.Hour).Format(time.RFC3339)

tests := []struct {
name string
useV1Public bool
statusCode int
responseBody string
shouldError bool
errorMsg string
name string
useV1Public bool
statusCode int
responseBody string
shouldError bool
errorMsg string
expectedBearer string // when empty, the assertion falls back to expectedToken
}{
{
name: "successful login with v1-public",
Expand Down Expand Up @@ -198,10 +199,29 @@ func TestRancherLogin(t *testing.T) {
shouldError: true,
errorMsg: "error unmarshaling",
},
{
name: "successful login with ext token response",
useV1Public: true,
statusCode: http.StatusCreated,
responseBody: fmt.Sprintf(`{
"apiVersion": "ext.cattle.io/v1",
"kind": "Token",
"metadata": {"name": "token-xyz"},
"spec": { "userID": "user-123" },
"status": {
"bearerToken": "ext/token-xyz:%s",
"expiresAt": "%s"
}
}`, expectedToken, expiresAt),
shouldError: false,
expectedBearer: "ext/token-xyz:" + expectedToken,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()

server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, http.MethodPost, r.Method)

Expand Down Expand Up @@ -252,7 +272,11 @@ func TestRancherLogin(t *testing.T) {
} else {
require.NoError(t, err)
assert.NotNil(t, token)
assert.Equal(t, expectedToken, token.Token)
want := tt.expectedBearer
if want == "" {
want = expectedToken
}
assert.Equal(t, want, token.BearerToken)
}
})
}
Expand Down Expand Up @@ -320,7 +344,7 @@ func TestOauthDeviceCodeAuth(t *testing.T) {

require.NoError(t, err)
assert.NotNil(t, token)
assert.Equal(t, "rancher-token-123", token.Token)
assert.Equal(t, "rancher-token-123", token.BearerToken)
}

func TestOauthAuthCodeAuth(t *testing.T) {
Expand Down
Loading
Loading