Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 4 additions & 7 deletions internal/smdclient/SMDclient.go
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ type SMDClient struct {
refreshLock sync.Mutex
clusterName string
smdClient *http.Client
tokenClient *http.Client
smdBaseURL string
tokenEndpoint string
accessToken string
Expand Down Expand Up @@ -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},

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should probably reuse the SMDClient transport here so TLS can be used for token handling. We can also reuse the timeout value too.

diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go
index 82ba48a..64319aa 100644
--- a/internal/smdclient/SMDclient.go
+++ b/internal/smdclient/SMDclient.go
@@ -94,7 +94,7 @@ func NewSMDClient(clusterName, baseurl, jwtURL, accessToken, certPath string, in
 		if err != nil {
 			return nil, fmt.Errorf("failed to read cert from path %s: %v", certPath, err)
 		}
-		certPool := x509.NewCertPool()
+		certPool = x509.NewCertPool()
 		certPool.AppendCertsFromPEM(cacert)
 	}
 
@@ -114,9 +114,12 @@ func NewSMDClient(clusterName, baseurl, jwtURL, accessToken, certPath string, in
 	}
 
 	client := &SMDClient{
-		clusterName:       clusterName,
-		smdClient:         c,
-		tokenClient:       &http.Client{Timeout: 10 * time.Second},
+		clusterName: clusterName,
+		smdClient:   c,
+		tokenClient: &http.Client{
+			Transport: c.Transport,
+			Timeout:   c.Timeout,
+		},
 		smdBaseURL:        baseurl,
 		tokenEndpoint:     jwtURL,
 		accessToken:       accessToken,

smdBaseURL: baseurl,
tokenEndpoint: jwtURL,
accessToken: accessToken,
Expand Down Expand Up @@ -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 {
Expand Down
1 change: 1 addition & 0 deletions internal/smdclient/SMDclient_stress_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
78 changes: 78 additions & 0 deletions internal/smdclient/SMDclient_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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) {
Expand Down
38 changes: 24 additions & 14 deletions internal/smdclient/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
}
Loading