src

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

credentials.go (3399B)


      1 package v4a
      2 
      3 import (
      4 	"context"
      5 	"crypto/ecdsa"
      6 	"fmt"
      7 	"sync"
      8 	"sync/atomic"
      9 	"time"
     10 
     11 	"github.com/aws/aws-sdk-go-v2/aws"
     12 	"github.com/aws/aws-sdk-go-v2/internal/sdk"
     13 )
     14 
     15 // Credentials is Context, ECDSA, and Optional Session Token that can be used
     16 // to sign requests using SigV4a
     17 type Credentials struct {
     18 	Context      string
     19 	PrivateKey   *ecdsa.PrivateKey
     20 	SessionToken string
     21 
     22 	// Time the credentials will expire.
     23 	CanExpire bool
     24 	Expires   time.Time
     25 }
     26 
     27 // Expired returns if the credentials have expired.
     28 func (v Credentials) Expired() bool {
     29 	if v.CanExpire {
     30 		return !v.Expires.After(sdk.NowTime())
     31 	}
     32 
     33 	return false
     34 }
     35 
     36 // HasKeys returns if the credentials keys are set.
     37 func (v Credentials) HasKeys() bool {
     38 	return len(v.Context) > 0 && v.PrivateKey != nil
     39 }
     40 
     41 // SymmetricCredentialAdaptor wraps a SigV4 AccessKey/SecretKey provider and adapts the credentials
     42 // to a ECDSA PrivateKey for signing with SiV4a
     43 type SymmetricCredentialAdaptor struct {
     44 	SymmetricProvider aws.CredentialsProvider
     45 
     46 	asymmetric atomic.Value
     47 	m          sync.Mutex
     48 }
     49 
     50 // Retrieve retrieves symmetric credentials from the underlying provider.
     51 func (s *SymmetricCredentialAdaptor) Retrieve(ctx context.Context) (aws.Credentials, error) {
     52 	symCreds, err := s.retrieveFromSymmetricProvider(ctx)
     53 	if err != nil {
     54 		return aws.Credentials{}, err
     55 	}
     56 
     57 	if asymCreds := s.getCreds(); asymCreds == nil {
     58 		return symCreds, nil
     59 	}
     60 
     61 	s.m.Lock()
     62 	defer s.m.Unlock()
     63 
     64 	asymCreds := s.getCreds()
     65 	if asymCreds == nil {
     66 		return symCreds, nil
     67 	}
     68 
     69 	// if the context does not match the access key id clear it
     70 	if asymCreds.Context != symCreds.AccessKeyID {
     71 		s.asymmetric.Store((*Credentials)(nil))
     72 	}
     73 
     74 	return symCreds, nil
     75 }
     76 
     77 // RetrievePrivateKey returns credentials suitable for SigV4a signing
     78 func (s *SymmetricCredentialAdaptor) RetrievePrivateKey(ctx context.Context) (Credentials, error) {
     79 	if asymCreds := s.getCreds(); asymCreds != nil {
     80 		return *asymCreds, nil
     81 	}
     82 
     83 	s.m.Lock()
     84 	defer s.m.Unlock()
     85 
     86 	if asymCreds := s.getCreds(); asymCreds != nil {
     87 		return *asymCreds, nil
     88 	}
     89 
     90 	symmetricCreds, err := s.retrieveFromSymmetricProvider(ctx)
     91 	if err != nil {
     92 		return Credentials{}, fmt.Errorf("failed to retrieve symmetric credentials: %v", err)
     93 	}
     94 
     95 	privateKey, err := deriveKeyFromAccessKeyPair(symmetricCreds.AccessKeyID, symmetricCreds.SecretAccessKey)
     96 	if err != nil {
     97 		return Credentials{}, fmt.Errorf("failed to derive assymetric key from credentials")
     98 	}
     99 
    100 	creds := Credentials{
    101 		Context:      symmetricCreds.AccessKeyID,
    102 		PrivateKey:   privateKey,
    103 		SessionToken: symmetricCreds.SessionToken,
    104 		CanExpire:    symmetricCreds.CanExpire,
    105 		Expires:      symmetricCreds.Expires,
    106 	}
    107 
    108 	s.asymmetric.Store(&creds)
    109 
    110 	return creds, nil
    111 }
    112 
    113 func (s *SymmetricCredentialAdaptor) getCreds() *Credentials {
    114 	v := s.asymmetric.Load()
    115 
    116 	if v == nil {
    117 		return nil
    118 	}
    119 
    120 	c := v.(*Credentials)
    121 	if c != nil && c.HasKeys() && !c.Expired() {
    122 		return c
    123 	}
    124 
    125 	return nil
    126 }
    127 
    128 func (s *SymmetricCredentialAdaptor) retrieveFromSymmetricProvider(ctx context.Context) (aws.Credentials, error) {
    129 	credentials, err := s.SymmetricProvider.Retrieve(ctx)
    130 	if err != nil {
    131 		return aws.Credentials{}, err
    132 	}
    133 
    134 	return credentials, nil
    135 }
    136 
    137 // CredentialsProvider is the interface for a provider to retrieve credentials
    138 // to sign requests with.
    139 type CredentialsProvider interface {
    140 	RetrievePrivateKey(context.Context) (Credentials, error)
    141 }