package oauth import ( "crypto/hmac" "crypto/rand" "crypto/sha1" "encoding/base64" "fmt" "slices" "strconv" "strings" "time" ) const ( oauthRealm = "realm" oauthConsumerKey = "oauth_consumer_key" oauthToken = "oauth_token" oauthNonce = "oauth_nonce" oauthTimestamp = "oauth_timestamp" oauthSignatureMethod = "oauth_signature_method" oauthVersion = "oauth_version" oauthSignature = "oauth_signature" ) const ( SignatureMethodHMACSHA1 = "HMAC-SHA1" ) const Version = "1.0" const separator = "&" type authHeader struct { method string baseURL string queryParams map[string]string tokens Tokens nonce string timestamp string signatureMethod string version string } func newAuthHeader( method string, baseURL string, queryParams map[string]string, tokens Tokens, ) authHeader { return authHeader{ method: method, baseURL: baseURL, queryParams: queryParams, tokens: tokens, nonce: rand.Text(), timestamp: strconv.Itoa(int(time.Now().Unix())), signatureMethod: SignatureMethodHMACSHA1, version: Version, } } func (ah *authHeader) String() string { signature, err := ah.buildSignature() if err != nil { return "" } subheaders := make(map[string]string, 8) subheaders[oauthRealm] = ah.baseURL subheaders[oauthVersion] = ah.version subheaders[oauthTimestamp] = ah.timestamp subheaders[oauthNonce] = ah.nonce subheaders[oauthConsumerKey] = ah.tokens.Consumer.Token subheaders[oauthToken] = ah.tokens.Access.Token subheaders[oauthSignatureMethod] = ah.signatureMethod subheaders[oauthSignature] = percentEncode(signature) parts := make([]string, 0, len(subheaders)) for k, v := range subheaders { parts = append(parts, fmt.Sprintf(`%s="%s"`, k, v)) } header := "OAuth " + strings.Join(parts, ", ") return header } func (ah *authHeader) buildSignature() (string, error) { baseString := ah.buildBaseString() baseString += ah.buildParameterString() signingKey := fmt.Appendf( nil, "%s&%s", percentEncode(ah.tokens.Consumer.Secret), percentEncode(ah.tokens.Access.Secret), ) hasher := hmac.New(sha1.New, signingKey) if _, err := hasher.Write([]byte(baseString)); err != nil { return "", err } encoded := base64.StdEncoding.EncodeToString(hasher.Sum(nil)) return encoded, nil } func (ah *authHeader) buildBaseString() string { baseString := strings.Join( []string{ strings.ToUpper(ah.method), percentEncode(ah.baseURL), }, separator, ) baseString += separator return baseString } func (ah *authHeader) buildParameterString() string { params := ah.queryParams if params == nil { params = make(map[string]string, 6) } params[oauthConsumerKey] = string(ah.tokens.Consumer.Token) params[oauthNonce] = string(ah.nonce) params[oauthSignatureMethod] = string(ah.signatureMethod) params[oauthTimestamp] = string(ah.timestamp) params[oauthToken] = string(ah.tokens.Access.Token) params[oauthVersion] = string(ah.version) keys := make([]string, 0, len(params)) for k := range params { keys = append(keys, k) } slices.Sort(keys) parts := make([]string, 0, len(keys)) for _, k := range keys { parts = append(parts, fmt.Sprintf(`%s=%s`, k, params[k])) } parameterString := strings.Join(parts, separator) parameterString = percentEncode(parameterString) return parameterString } func percentEncode(s string) string { var b strings.Builder for _, c := range s { if isUnreservedCharacter(c) { fmt.Fprintf(&b, "%%%02X", c) continue } b.WriteRune(c) } return b.String() } // isUnreservedCharacter returns true if the byte belongs in the range of the // reserved character list as in https://en.wikipedia.org/wiki/Percent-encoding func isUnreservedCharacter(c rune) bool { if (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '-' || c == '.' || c == '_' || c == '~' { return false } return true }