diff options
Diffstat (limited to 'src/oauth/auth_header.go')
| -rw-r--r-- | src/oauth/auth_header.go | 181 |
1 files changed, 181 insertions, 0 deletions
diff --git a/src/oauth/auth_header.go b/src/oauth/auth_header.go new file mode 100644 index 0000000..bf30812 --- /dev/null +++ b/src/oauth/auth_header.go @@ -0,0 +1,181 @@ +package oauth + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha1" + "encoding/base64" + "fmt" + "sort" + "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, _ := ah.buildSignature() // TODO: handle ignored err + + 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 string(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) + } + sort.Strings(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 +} |
