token_provider.go (7164B)
1 package imds 2 3 import ( 4 "context" 5 "errors" 6 "fmt" 7 "net/http" 8 "sync" 9 "sync/atomic" 10 "time" 11 12 "github.com/aws/aws-sdk-go-v2/aws" 13 "github.com/aws/smithy-go" 14 "github.com/aws/smithy-go/logging" 15 16 "github.com/aws/smithy-go/middleware" 17 smithyhttp "github.com/aws/smithy-go/transport/http" 18 ) 19 20 const ( 21 // Headers for Token and TTL 22 tokenHeader = "x-aws-ec2-metadata-token" 23 defaultTokenTTL = 5 * time.Minute 24 ) 25 26 type tokenProvider struct { 27 client *Client 28 tokenTTL time.Duration 29 30 token *apiToken 31 tokenMux sync.RWMutex 32 33 disabled uint32 // Atomic updated 34 } 35 36 func newTokenProvider(client *Client, ttl time.Duration) *tokenProvider { 37 return &tokenProvider{ 38 client: client, 39 tokenTTL: ttl, 40 } 41 } 42 43 // apiToken provides the API token used by all operation calls for th EC2 44 // Instance metadata service. 45 type apiToken struct { 46 token string 47 expires time.Time 48 } 49 50 var timeNow = time.Now 51 52 // Expired returns if the token is expired. 53 func (t *apiToken) Expired() bool { 54 // Calling Round(0) on the current time will truncate the monotonic reading only. Ensures credential expiry 55 // time is always based on reported wall-clock time. 56 return timeNow().Round(0).After(t.expires) 57 } 58 59 func (t *tokenProvider) ID() string { return "APITokenProvider" } 60 61 // HandleFinalize is the finalize stack middleware, that if the token provider is 62 // enabled, will attempt to add the cached API token to the request. If the API 63 // token is not cached, it will be retrieved in a separate API call, getToken. 64 // 65 // For retry attempts, handler must be added after attempt retryer. 66 // 67 // If request for getToken fails the token provider may be disabled from future 68 // requests, depending on the response status code. 69 func (t *tokenProvider) HandleFinalize( 70 ctx context.Context, input middleware.FinalizeInput, next middleware.FinalizeHandler, 71 ) ( 72 out middleware.FinalizeOutput, metadata middleware.Metadata, err error, 73 ) { 74 if t.fallbackEnabled() && !t.enabled() { 75 // short-circuits to insecure data flow if token provider is disabled. 76 return next.HandleFinalize(ctx, input) 77 } 78 79 req, ok := input.Request.(*smithyhttp.Request) 80 if !ok { 81 return out, metadata, fmt.Errorf("unexpected transport request type %T", input.Request) 82 } 83 84 tok, err := t.getToken(ctx) 85 if err != nil { 86 // If the error allows the token to downgrade to insecure flow allow that. 87 var bypassErr *bypassTokenRetrievalError 88 if errors.As(err, &bypassErr) { 89 return next.HandleFinalize(ctx, input) 90 } 91 92 return out, metadata, fmt.Errorf("failed to get API token, %w", err) 93 } 94 95 req.Header.Set(tokenHeader, tok.token) 96 97 return next.HandleFinalize(ctx, input) 98 } 99 100 // HandleDeserialize is the deserialize stack middleware for determining if the 101 // operation the token provider is decorating failed because of a 401 102 // unauthorized status code. If the operation failed for that reason the token 103 // provider needs to be re-enabled so that it can start adding the API token to 104 // operation calls. 105 func (t *tokenProvider) HandleDeserialize( 106 ctx context.Context, input middleware.DeserializeInput, next middleware.DeserializeHandler, 107 ) ( 108 out middleware.DeserializeOutput, metadata middleware.Metadata, err error, 109 ) { 110 out, metadata, err = next.HandleDeserialize(ctx, input) 111 if err == nil { 112 return out, metadata, err 113 } 114 115 resp, ok := out.RawResponse.(*smithyhttp.Response) 116 if !ok { 117 return out, metadata, fmt.Errorf("expect HTTP transport, got %T", out.RawResponse) 118 } 119 120 if resp.StatusCode == http.StatusUnauthorized { // unauthorized 121 t.enable() 122 err = &retryableError{Err: err, isRetryable: true} 123 } 124 125 return out, metadata, err 126 } 127 128 func (t *tokenProvider) getToken(ctx context.Context) (tok *apiToken, err error) { 129 if t.fallbackEnabled() && !t.enabled() { 130 return nil, &bypassTokenRetrievalError{ 131 Err: fmt.Errorf("cannot get API token, provider disabled"), 132 } 133 } 134 135 t.tokenMux.RLock() 136 tok = t.token 137 t.tokenMux.RUnlock() 138 139 if tok != nil && !tok.Expired() { 140 return tok, nil 141 } 142 143 tok, err = t.updateToken(ctx) 144 if err != nil { 145 return nil, err 146 } 147 148 return tok, nil 149 } 150 151 func (t *tokenProvider) updateToken(ctx context.Context) (*apiToken, error) { 152 t.tokenMux.Lock() 153 defer t.tokenMux.Unlock() 154 155 // Prevent multiple requests to update retrieving the token. 156 if t.token != nil && !t.token.Expired() { 157 tok := t.token 158 return tok, nil 159 } 160 161 result, err := t.client.getToken(ctx, &getTokenInput{ 162 TokenTTL: t.tokenTTL, 163 }) 164 if err != nil { 165 var statusErr interface{ HTTPStatusCode() int } 166 if errors.As(err, &statusErr) { 167 switch statusErr.HTTPStatusCode() { 168 // Disable future get token if failed because of 403, 404, or 405 169 case http.StatusForbidden, 170 http.StatusNotFound, 171 http.StatusMethodNotAllowed: 172 173 if t.fallbackEnabled() { 174 logger := middleware.GetLogger(ctx) 175 logger.Logf(logging.Warn, "falling back to IMDSv1: %v", err) 176 t.disable() 177 } 178 179 // 400 errors are terminal, and need to be upstreamed 180 case http.StatusBadRequest: 181 return nil, err 182 } 183 } 184 185 // Disable if request send failed or timed out getting response 186 var re *smithyhttp.RequestSendError 187 var ce *smithy.CanceledError 188 if errors.As(err, &re) || errors.As(err, &ce) { 189 atomic.StoreUint32(&t.disabled, 1) 190 } 191 192 if !t.fallbackEnabled() { 193 // NOTE: getToken() is an implementation detail of some outer operation 194 // (e.g. GetMetadata). It has its own retries that have already been exhausted. 195 // Mark the underlying error as a terminal error. 196 err = &retryableError{Err: err, isRetryable: false} 197 return nil, err 198 } 199 200 // Token couldn't be retrieved, fallback to IMDSv1 insecure flow for this request 201 // and allow the request to proceed. Future requests _may_ re-attempt fetching a 202 // token if not disabled. 203 return nil, &bypassTokenRetrievalError{Err: err} 204 } 205 206 tok := &apiToken{ 207 token: result.Token, 208 expires: timeNow().Add(result.TokenTTL), 209 } 210 t.token = tok 211 212 return tok, nil 213 } 214 215 // enabled returns if the token provider is current enabled or not. 216 func (t *tokenProvider) enabled() bool { 217 return atomic.LoadUint32(&t.disabled) == 0 218 } 219 220 // fallbackEnabled returns false if EnableFallback is [aws.FalseTernary], true otherwise 221 func (t *tokenProvider) fallbackEnabled() bool { 222 switch t.client.options.EnableFallback { 223 case aws.FalseTernary: 224 return false 225 default: 226 return true 227 } 228 } 229 230 // disable disables the token provider and it will no longer attempt to inject 231 // the token, nor request updates. 232 func (t *tokenProvider) disable() { 233 atomic.StoreUint32(&t.disabled, 1) 234 } 235 236 // enable enables the token provide to start refreshing tokens, and adding them 237 // to the pending request. 238 func (t *tokenProvider) enable() { 239 t.tokenMux.Lock() 240 t.token = nil 241 t.tokenMux.Unlock() 242 atomic.StoreUint32(&t.disabled, 0) 243 } 244 245 type bypassTokenRetrievalError struct { 246 Err error 247 } 248 249 func (e *bypassTokenRetrievalError) Error() string { 250 return fmt.Sprintf("bypass token retrieval, %v", e.Err) 251 } 252 253 func (e *bypassTokenRetrievalError) Unwrap() error { return e.Err } 254 255 type retryableError struct { 256 Err error 257 isRetryable bool 258 } 259 260 func (e *retryableError) RetryableError() bool { return e.isRetryable } 261 262 func (e *retryableError) Error() string { return e.Err.Error() }