From 8ea5ab3b1382c9edda5922d00ded77a5504a257f Mon Sep 17 00:00:00 2001 From: Alex Lovell-Troy Date: Wed, 16 Sep 2026 17:26:58 -0400 Subject: [PATCH] fix: enhance token refresh handling and improve error reporting Signed-off-by: Alex Lovell-Troy --- internal/smdclient/SMDclient.go | 11 ++- internal/smdclient/SMDclient_stress_test.go | 1 + internal/smdclient/SMDclient_test.go | 78 +++++++++++++++++++++ internal/smdclient/oidc.go | 38 ++++++---- 4 files changed, 107 insertions(+), 21 deletions(-) diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go index eb1bd91e..82ba48a0 100644 --- a/internal/smdclient/SMDclient.go +++ b/internal/smdclient/SMDclient.go @@ -48,6 +48,7 @@ type SMDClient struct { refreshLock sync.Mutex clusterName string smdClient *http.Client + tokenClient *http.Client smdBaseURL string tokenEndpoint string accessToken string @@ -115,6 +116,7 @@ func NewSMDClient(clusterName, baseurl, jwtURL, accessToken, certPath string, in client := &SMDClient{ clusterName: clusterName, smdClient: c, + tokenClient: &http.Client{Timeout: 10 * time.Second}, smdBaseURL: baseurl, tokenEndpoint: jwtURL, accessToken: accessToken, @@ -197,13 +199,8 @@ func (s *SMDClient) getSMD(ep string, smd any) error { if !freshToken { log.Info().Msg("Fetching new JWT and retrying...") // Try to refresh the token and retry once - if err2 := s.refreshTokenIfCurrent(usedToken); err2 != nil { - // If token refresh fails, refresh will attempt again. - // While effectively we could ignore the error, it helps - // to see why the failure is occurring in case the error - // is unusual (RefreshToken() has a few different failure - // modes). - log.Debug().Err(err).Msg("failed to refresh token") + if err := s.refreshTokenIfCurrent(usedToken); err != nil { + return fmt.Errorf("refreshing rejected SMD access token: %w", err) } freshToken = true } else { diff --git a/internal/smdclient/SMDclient_stress_test.go b/internal/smdclient/SMDclient_stress_test.go index b8154bba..66557052 100644 --- a/internal/smdclient/SMDclient_stress_test.go +++ b/internal/smdclient/SMDclient_stress_test.go @@ -84,6 +84,7 @@ func TestStressConcurrentGetSMDCoalescesTokenRefresh10K(t *testing.T) { client := &SMDClient{ smdClient: &http.Client{Transport: tokenAwareRoundTripper{}}, + tokenClient: tokenServer.Client(), smdBaseURL: "http://smd.example", tokenEndpoint: tokenServer.URL, accessToken: "stale-token", diff --git a/internal/smdclient/SMDclient_test.go b/internal/smdclient/SMDclient_test.go index f93d06a1..a6fe1f3a 100644 --- a/internal/smdclient/SMDclient_test.go +++ b/internal/smdclient/SMDclient_test.go @@ -165,6 +165,7 @@ func TestConcurrentGetSMDCoalescesTokenRefresh(t *testing.T) { client := &SMDClient{ smdClient: smdServer.Client(), + tokenClient: tokenServer.Client(), smdBaseURL: smdServer.URL, tokenEndpoint: tokenServer.URL, accessToken: "stale-token", @@ -205,6 +206,83 @@ func TestConcurrentGetSMDCoalescesTokenRefresh(t *testing.T) { } } +func TestRefreshTokenFailurePreservesAccessToken(t *testing.T) { + tests := []struct { + name string + statusCode int + body string + }{ + { + name: "non-2xx response", + statusCode: http.StatusInternalServerError, + body: `{"access_token":"replacement-token"}`, + }, + { + name: "empty access token", + statusCode: http.StatusOK, + body: `{"access_token":" \t "}`, + }, + { + name: "malformed JSON", + statusCode: http.StatusOK, + body: `{"access_token":`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(tt.statusCode) + _, _ = w.Write([]byte(tt.body)) + })) + defer tokenServer.Close() + + client := &SMDClient{ + tokenEndpoint: tokenServer.URL, + tokenClient: tokenServer.Client(), + accessToken: "previous-token", + } + + err := client.RefreshToken() + + require.Error(t, err) + assert.Equal(t, "previous-token", client.currentAccessToken()) + }) + } +} + +func TestGetSMDReturnsTokenRefreshFailureWithoutRetry(t *testing.T) { + var smdRequests atomic.Int64 + smdServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + smdRequests.Add(1) + w.WriteHeader(http.StatusUnauthorized) + })) + defer smdServer.Close() + + var tokenRequests atomic.Int64 + tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tokenRequests.Add(1) + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer tokenServer.Close() + + client := &SMDClient{ + smdClient: smdServer.Client(), + tokenClient: tokenServer.Client(), + smdBaseURL: smdServer.URL, + tokenEndpoint: tokenServer.URL, + accessToken: "previous-token", + } + + var response map[string]string + err := client.getSMD("/component", &response) + + require.ErrorContains(t, err, "refreshing rejected SMD access token") + assert.Equal(t, int64(1), smdRequests.Load()) + assert.Equal(t, int64(1), tokenRequests.Load()) + assert.Equal(t, "previous-token", client.currentAccessToken()) +} + func TestComponentInformationUsesCache(t *testing.T) { var perNodeComponentRequests atomic.Int64 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/internal/smdclient/oidc.go b/internal/smdclient/oidc.go index b54e0820..4c04e7fd 100644 --- a/internal/smdclient/oidc.go +++ b/internal/smdclient/oidc.go @@ -10,15 +10,16 @@ import ( "fmt" "io" "net/http" + "strings" "time" ) // Structure of a token reponse from OIDC server type oidcTokenData struct { - Access_token string `json:"access_token" yaml:"access_token"` - Expires_in int `json:"expires_in" yaml:"expires_in"` - Scope string `json:"scope" yaml:"scope"` - Token_type string `json:"token_type" yaml:"token_type"` + AccessToken string `json:"access_token" yaml:"access_token"` + ExpiresIn int `json:"expires_in" yaml:"expires_in"` + Scope string `json:"scope" yaml:"scope"` + TokenType string `json:"token_type" yaml:"token_type"` } // Refresh the cached access token, using the provided JWT server @@ -64,31 +65,40 @@ func (s *SMDClient) refreshTokenIfCurrent(rejectedToken string) error { } func (s *SMDClient) refreshTokenWithContext(ctx context.Context) error { + if s.tokenClient == nil { + return fmt.Errorf("token HTTP client is not configured") + } + // Request new token from OIDC server using the provided context. - req, err := http.NewRequestWithContext(ctx, "GET", s.tokenEndpoint, nil) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, s.tokenEndpoint, nil) if err != nil { - return err - } - if s.smdClient == nil { - return fmt.Errorf("SMD HTTP client was nil (was NewSMDClient() run?)") + return fmt.Errorf("creating token request: %w", err) } - r, err := s.smdClient.Do(req) + r, err := s.tokenClient.Do(req) if err != nil { - return err + return fmt.Errorf("requesting token: %w", err) } defer r.Body.Close() + if r.StatusCode < http.StatusOK || r.StatusCode >= http.StatusMultipleChoices { + return fmt.Errorf("token endpoint returned HTTP %d", r.StatusCode) + } body, err := io.ReadAll(r.Body) if err != nil { - return err + return fmt.Errorf("reading token response: %w", err) } // Decode server's response to the expected structure var tokenResp oidcTokenData if err = json.Unmarshal(body, &tokenResp); err != nil { - return err + return fmt.Errorf("decoding token response: %w", err) + } + token := strings.TrimSpace(tokenResp.AccessToken) + if token == "" { + return fmt.Errorf("token response contains an empty access token") } + // Store the JWT safely. s.accessTokenMutex.Lock() - s.accessToken = tokenResp.Access_token + s.accessToken = token s.accessTokenMutex.Unlock() return nil }