diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go index c2c62fdf..712d36b3 100644 --- a/internal/smdclient/SMDclient.go +++ b/internal/smdclient/SMDclient.go @@ -41,6 +41,7 @@ type SMDClientInterface interface { // SMDClient is a client for SMD type SMDClient struct { + refreshLock sync.Mutex clusterName string smdClient *http.Client smdBaseURL string diff --git a/internal/smdclient/oidc.go b/internal/smdclient/oidc.go index dfc77105..bbb7e07b 100644 --- a/internal/smdclient/oidc.go +++ b/internal/smdclient/oidc.go @@ -1,9 +1,12 @@ package smdclient import ( + "context" "encoding/json" + "fmt" "io" "net/http" + "time" ) // Structure of a token reponse from OIDC server @@ -19,23 +22,53 @@ type oidcTokenData struct { // authorization grant. Support for said grant should probably be implemented // at some point. func (s *SMDClient) RefreshToken() error { - s.accessTokenMutex.Lock() - defer s.accessTokenMutex.Unlock() - return s.refreshTokenLocked() + // Serialize refresh to avoid concurrent token fetches. + s.refreshLock.Lock() + defer s.refreshLock.Unlock() + ctx, cancel := context.WithTimeout(context.Background(), defaultRefreshTimeout) + defer cancel() + return s.refreshTokenWithContext(ctx) } +const defaultRefreshTimeout = 10 * time.Second + func (s *SMDClient) refreshTokenIfCurrent(rejectedToken string) error { + // Fast path: if token already different, nothing to do. + s.accessTokenMutex.Lock() + if s.accessToken != rejectedToken { + s.accessTokenMutex.Unlock() + return nil + } + s.accessTokenMutex.Unlock() + + // Serialize refresh to avoid concurrent token fetches. + s.refreshLock.Lock() + defer s.refreshLock.Unlock() + + // Re-check token after acquiring lock (it may have been refreshed by another goroutine). s.accessTokenMutex.Lock() - defer s.accessTokenMutex.Unlock() if s.accessToken != rejectedToken { + s.accessTokenMutex.Unlock() return nil } - return s.refreshTokenLocked() + s.accessTokenMutex.Unlock() + + // Acquire new token with timeout. + ctx, cancel := context.WithTimeout(context.Background(), defaultRefreshTimeout) + defer cancel() + return s.refreshTokenWithContext(ctx) } -func (s *SMDClient) refreshTokenLocked() error { - // Request new token from OIDC server - r, err := http.Get(s.tokenEndpoint) +func (s *SMDClient) refreshTokenWithContext(ctx context.Context) error { + // Request new token from OIDC server using the provided context. + req, err := http.NewRequestWithContext(ctx, "GET", s.tokenEndpoint, nil) + if err != nil { + return err + } + if s.smdClient == nil { + return fmt.Errorf("SMD HTTP client was nil (was NewSMDClient() run?)") + } + r, err := s.smdClient.Do(req) if err != nil { return err } @@ -49,7 +82,9 @@ func (s *SMDClient) refreshTokenLocked() error { if err = json.Unmarshal(body, &tokenResp); err != nil { return err } - // Extract and store the JWT itself + // Store the JWT safely. + s.accessTokenMutex.Lock() s.accessToken = tokenResp.Access_token + s.accessTokenMutex.Unlock() return nil }