src

Go monorepo.
git clone git://code.dwrz.net/src
Log | Files | Refs

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() }