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 }