src

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

sso_cached_token.go (5838B)


      1 package ssocreds
      2 
      3 import (
      4 	"crypto/sha1"
      5 	"encoding/hex"
      6 	"encoding/json"
      7 	"fmt"
      8 	"os"
      9 	"path/filepath"
     10 	"strconv"
     11 	"strings"
     12 	"time"
     13 
     14 	"github.com/aws/aws-sdk-go-v2/internal/sdk"
     15 	"github.com/aws/aws-sdk-go-v2/internal/shareddefaults"
     16 )
     17 
     18 var osUserHomeDur = shareddefaults.UserHomeDir
     19 
     20 // StandardCachedTokenFilepath returns the filepath for the cached SSO token file, or
     21 // error if unable get derive the path. Key that will be used to compute a SHA1
     22 // value that is hex encoded.
     23 //
     24 // Derives the filepath using the Key as:
     25 //
     26 //	~/.aws/sso/cache/<sha1-hex-encoded-key>.json
     27 func StandardCachedTokenFilepath(key string) (string, error) {
     28 	homeDir := osUserHomeDur()
     29 	if len(homeDir) == 0 {
     30 		return "", fmt.Errorf("unable to get USER's home directory for cached token")
     31 	}
     32 	hash := sha1.New()
     33 	if _, err := hash.Write([]byte(key)); err != nil {
     34 		return "", fmt.Errorf("unable to compute cached token filepath key SHA1 hash, %w", err)
     35 	}
     36 
     37 	cacheFilename := strings.ToLower(hex.EncodeToString(hash.Sum(nil))) + ".json"
     38 
     39 	return filepath.Join(homeDir, ".aws", "sso", "cache", cacheFilename), nil
     40 }
     41 
     42 type tokenKnownFields struct {
     43 	AccessToken string   `json:"accessToken,omitempty"`
     44 	ExpiresAt   *rfc3339 `json:"expiresAt,omitempty"`
     45 
     46 	RefreshToken string `json:"refreshToken,omitempty"`
     47 	ClientID     string `json:"clientId,omitempty"`
     48 	ClientSecret string `json:"clientSecret,omitempty"`
     49 }
     50 
     51 type token struct {
     52 	tokenKnownFields
     53 	UnknownFields map[string]interface{} `json:"-"`
     54 }
     55 
     56 func (t token) MarshalJSON() ([]byte, error) {
     57 	fields := map[string]interface{}{}
     58 
     59 	setTokenFieldString(fields, "accessToken", t.AccessToken)
     60 	setTokenFieldRFC3339(fields, "expiresAt", t.ExpiresAt)
     61 
     62 	setTokenFieldString(fields, "refreshToken", t.RefreshToken)
     63 	setTokenFieldString(fields, "clientId", t.ClientID)
     64 	setTokenFieldString(fields, "clientSecret", t.ClientSecret)
     65 
     66 	for k, v := range t.UnknownFields {
     67 		if _, ok := fields[k]; ok {
     68 			return nil, fmt.Errorf("unknown token field %v, duplicates known field", k)
     69 		}
     70 		fields[k] = v
     71 	}
     72 
     73 	return json.Marshal(fields)
     74 }
     75 
     76 func setTokenFieldString(fields map[string]interface{}, key, value string) {
     77 	if value == "" {
     78 		return
     79 	}
     80 	fields[key] = value
     81 }
     82 func setTokenFieldRFC3339(fields map[string]interface{}, key string, value *rfc3339) {
     83 	if value == nil {
     84 		return
     85 	}
     86 	fields[key] = value
     87 }
     88 
     89 func (t *token) UnmarshalJSON(b []byte) error {
     90 	var fields map[string]interface{}
     91 	if err := json.Unmarshal(b, &fields); err != nil {
     92 		return nil
     93 	}
     94 
     95 	t.UnknownFields = map[string]interface{}{}
     96 
     97 	for k, v := range fields {
     98 		var err error
     99 		switch k {
    100 		case "accessToken":
    101 			err = getTokenFieldString(v, &t.AccessToken)
    102 		case "expiresAt":
    103 			err = getTokenFieldRFC3339(v, &t.ExpiresAt)
    104 		case "refreshToken":
    105 			err = getTokenFieldString(v, &t.RefreshToken)
    106 		case "clientId":
    107 			err = getTokenFieldString(v, &t.ClientID)
    108 		case "clientSecret":
    109 			err = getTokenFieldString(v, &t.ClientSecret)
    110 		default:
    111 			t.UnknownFields[k] = v
    112 		}
    113 
    114 		if err != nil {
    115 			return fmt.Errorf("field %q, %w", k, err)
    116 		}
    117 	}
    118 
    119 	return nil
    120 }
    121 
    122 func getTokenFieldString(v interface{}, value *string) error {
    123 	var ok bool
    124 	*value, ok = v.(string)
    125 	if !ok {
    126 		return fmt.Errorf("expect value to be string, got %T", v)
    127 	}
    128 	return nil
    129 }
    130 
    131 func getTokenFieldRFC3339(v interface{}, value **rfc3339) error {
    132 	var stringValue string
    133 	if err := getTokenFieldString(v, &stringValue); err != nil {
    134 		return err
    135 	}
    136 
    137 	timeValue, err := parseRFC3339(stringValue)
    138 	if err != nil {
    139 		return err
    140 	}
    141 
    142 	*value = &timeValue
    143 	return nil
    144 }
    145 
    146 func loadCachedToken(filename string) (token, error) {
    147 	fileBytes, err := os.ReadFile(filename)
    148 	if err != nil {
    149 		return token{}, fmt.Errorf("failed to read cached SSO token file, %w", err)
    150 	}
    151 
    152 	var t token
    153 	if err := json.Unmarshal(fileBytes, &t); err != nil {
    154 		return token{}, fmt.Errorf("failed to parse cached SSO token file, %w", err)
    155 	}
    156 
    157 	if len(t.AccessToken) == 0 || t.ExpiresAt == nil || time.Time(*t.ExpiresAt).IsZero() {
    158 		return token{}, fmt.Errorf(
    159 			"cached SSO token must contain accessToken and expiresAt fields")
    160 	}
    161 
    162 	return t, nil
    163 }
    164 
    165 func storeCachedToken(filename string, t token, fileMode os.FileMode) (err error) {
    166 	tmpFilename := filename + ".tmp-" + strconv.FormatInt(sdk.NowTime().UnixNano(), 10)
    167 	if err := writeCacheFile(tmpFilename, fileMode, t); err != nil {
    168 		return err
    169 	}
    170 
    171 	if err := os.Rename(tmpFilename, filename); err != nil {
    172 		return fmt.Errorf("failed to replace old cached SSO token file, %w", err)
    173 	}
    174 
    175 	return nil
    176 }
    177 
    178 func writeCacheFile(filename string, fileMode os.FileMode, t token) (err error) {
    179 	var f *os.File
    180 	f, err = os.OpenFile(filename, os.O_CREATE|os.O_TRUNC|os.O_RDWR, fileMode)
    181 	if err != nil {
    182 		return fmt.Errorf("failed to create cached SSO token file %w", err)
    183 	}
    184 
    185 	defer func() {
    186 		closeErr := f.Close()
    187 		if err == nil && closeErr != nil {
    188 			err = fmt.Errorf("failed to close cached SSO token file, %w", closeErr)
    189 		}
    190 	}()
    191 
    192 	encoder := json.NewEncoder(f)
    193 
    194 	if err = encoder.Encode(t); err != nil {
    195 		return fmt.Errorf("failed to serialize cached SSO token, %w", err)
    196 	}
    197 
    198 	return nil
    199 }
    200 
    201 type rfc3339 time.Time
    202 
    203 func parseRFC3339(v string) (rfc3339, error) {
    204 	parsed, err := time.Parse(time.RFC3339, v)
    205 	if err != nil {
    206 		return rfc3339{}, fmt.Errorf("expected RFC3339 timestamp: %w", err)
    207 	}
    208 
    209 	return rfc3339(parsed), nil
    210 }
    211 
    212 func (r *rfc3339) UnmarshalJSON(bytes []byte) (err error) {
    213 	var value string
    214 
    215 	// Use JSON unmarshal to unescape the quoted value making use of JSON's
    216 	// unquoting rules.
    217 	if err = json.Unmarshal(bytes, &value); err != nil {
    218 		return err
    219 	}
    220 
    221 	*r, err = parseRFC3339(value)
    222 
    223 	return nil
    224 }
    225 
    226 func (r *rfc3339) MarshalJSON() ([]byte, error) {
    227 	value := time.Time(*r).UTC().Format(time.RFC3339)
    228 
    229 	// Use JSON unmarshal to unescape the quoted value making use of JSON's
    230 	// quoting rules.
    231 	return json.Marshal(value)
    232 }