summaryrefslogtreecommitdiff
path: root/src/oauth/auth_header.go
diff options
context:
space:
mode:
authorDJ O'Leary <dijitol@proton.me>2026-09-03 02:35:45 +0200
committerDJ O'Leary <dijitol@proton.me>2026-09-03 02:35:45 +0200
commit1fe1e04c3b413109bd3d7fddea975513c455b098 (patch)
treeba2198542be5305b9c2eccac5c24503fbf306aaa /src/oauth/auth_header.go
parentf5440cff6298e4a29c5af42ef38f690b10c853bd (diff)
feat(oauth): create oauth1 client
Diffstat (limited to 'src/oauth/auth_header.go')
-rw-r--r--src/oauth/auth_header.go181
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
+}