From fc3eef495a42d578da3da1b2751726ec5fce40b8 Mon Sep 17 00:00:00 2001 From: Will Andrews Date: Fri, 25 Apr 2025 20:59:05 +0100 Subject: [PATCH] some improvements to the authflow when creating a new status --- auth_handlers.go | 64 ------------------------------------------ go.mod | 3 +- go.sum | 2 -- oauth/service.go | 4 +++ status.go | 73 ++++++++++++++++++++++++++++-------------------- 5 files changed, 47 insertions(+), 99 deletions(-) diff --git a/auth_handlers.go b/auth_handlers.go index 140f319..8b10fac 100644 --- a/auth_handlers.go +++ b/auth_handlers.go @@ -1,20 +1,13 @@ package statusphere import ( - "crypto/sha256" _ "embed" - "encoding/base64" - "encoding/json" "fmt" "log/slog" "net/http" "net/url" - "time" - "github.com/golang-jwt/jwt" - "github.com/google/uuid" "github.com/gorilla/sessions" - "github.com/lestrrat-go/jwx/v2/jwk" "github.com/willdot/statusphere-go/oauth" ) @@ -195,60 +188,3 @@ func (s *Server) HandleLogOut(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/", http.StatusFound) } - -func pdsDpopJwt(method, url, iss, accessToken, nonce string, privateJwk jwk.Key) (string, error) { - pubJwk, err := privateJwk.PublicKey() - if err != nil { - return "", err - } - - b, err := json.Marshal(pubJwk) - if err != nil { - return "", err - } - - var pubMap map[string]any - if err := json.Unmarshal(b, &pubMap); err != nil { - return "", err - } - - now := time.Now().Unix() - - claims := jwt.MapClaims{ - "iss": iss, - "iat": now, - "exp": now + 30, - "jti": uuid.NewString(), - "htm": method, - "htu": url, - "ath": generateCodeChallenge(accessToken), - } - - if nonce != "" { - claims["nonce"] = nonce - } - - token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) - token.Header["typ"] = "dpop+jwt" - token.Header["alg"] = "ES256" - token.Header["jwk"] = pubMap - - var rawKey any - if err := privateJwk.Raw(&rawKey); err != nil { - return "", err - } - - tokenString, err := token.SignedString(rawKey) - if err != nil { - return "", fmt.Errorf("failed to sign token: %w", err) - } - - return tokenString, nil -} - -func generateCodeChallenge(pkceVerifier string) string { - h := sha256.New() - h.Write([]byte(pkceVerifier)) - hash := h.Sum(nil) - return base64.RawURLEncoding.EncodeToString(hash) -} diff --git a/go.mod b/go.mod index 1e8dcbf..c6fc000 100644 --- a/go.mod +++ b/go.mod @@ -8,8 +8,6 @@ require ( github.com/avast/retry-go/v4 v4.6.1 github.com/bluesky-social/jetstream v0.0.0-20250414024304-d17bd81a945e github.com/glebarez/go-sqlite v1.22.0 - github.com/golang-jwt/jwt v3.2.2+incompatible - github.com/google/uuid v1.6.0 github.com/gorilla/sessions v1.4.0 github.com/haileyok/atproto-oauth-golang v0.0.2 github.com/joho/godotenv v1.5.1 @@ -29,6 +27,7 @@ require ( github.com/goccy/go-json v0.10.3 // indirect github.com/gogo/protobuf v1.3.2 // indirect github.com/golang-jwt/jwt/v5 v5.2.1 // indirect + github.com/google/uuid v1.6.0 // indirect github.com/gorilla/securecookie v1.1.2 // indirect github.com/gorilla/websocket v1.5.1 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect diff --git a/go.sum b/go.sum index dea16ba..20d6a9d 100644 --- a/go.sum +++ b/go.sum @@ -38,8 +38,6 @@ github.com/goccy/go-json v0.10.3 h1:KZ5WoDbxAIgm2HNbYckL0se1fHD6rz5j4ywS6ebzDqA= github.com/goccy/go-json v0.10.3/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= -github.com/golang-jwt/jwt v3.2.2+incompatible h1:IfV12K8xAKAnZqdXVzCZ+TOjboZ2keLg81eXfW3O+oY= -github.com/golang-jwt/jwt v3.2.2+incompatible/go.mod h1:8pz2t5EyA70fFQQSrl6XZXzqecmYZeUEB8OUGHkxJ+I= github.com/golang-jwt/jwt/v5 v5.2.1 h1:OuVbFODueb089Lh128TAcimifWaLhJwVflnrgM17wHk= github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= diff --git a/oauth/service.go b/oauth/service.go index f376978..58fa101 100644 --- a/oauth/service.go +++ b/oauth/service.go @@ -244,6 +244,10 @@ func (s *Service) PublicKey() []byte { return s.jwks.public } +func (s *Service) PdsDpopJwt(method, url string, session Session, privateKey jwk.Key) (string, error) { + return atoauth.PdsDpopJwt(method, url, session.AuthserverIss, session.AccessToken, session.DpopPdsNonce, privateKey) +} + func (s *Service) makeOAuthRequest(ctx context.Context, did, handle string, dpopPrivateKey jwk.Key) (*atoauth.SendParAuthResponse, *atoauth.OauthAuthorizationMetadata, string, error) { service, err := s.resolveService(ctx, did) if err != nil { diff --git a/status.go b/status.go index c50b157..ad8e4e7 100644 --- a/status.go +++ b/status.go @@ -27,15 +27,12 @@ type XRPCError struct { } type CreateRecordResp struct { - URI string `json:"uri"` + URI string `json:"uri"` + ErrStr string `json:"error"` + Message string `json:"message"` } func (s *Server) CreateNewStatus(ctx context.Context, oauthsession oauth.Session, status string, createdAt time.Time) (string, error) { - privateJwk, err := oauthsession.CreatePrivateKey() - if err != nil { - return "", fmt.Errorf("create private jwk: %w", err) - } - bodyReq := map[string]any{ "repo": oauthsession.Did, "collection": "xyz.statusphere.status", @@ -50,25 +47,31 @@ func (s *Server) CreateNewStatus(ctx context.Context, oauthsession oauth.Session return "", fmt.Errorf("marshal update message request body: %w", err) } - // TODO: redo this loop business - for range 2 { - r := bytes.NewReader(bodyB) - url := fmt.Sprintf("%s/xrpc/com.atproto.repo.createRecord", oauthsession.PdsUrl) - request, err := http.NewRequestWithContext(ctx, "POST", url, r) - if err != nil { - return "", fmt.Errorf("create http request: %w", err) - } + r := bytes.NewReader(bodyB) + url := fmt.Sprintf("%s/xrpc/com.atproto.repo.createRecord", oauthsession.PdsUrl) + request, err := http.NewRequestWithContext(ctx, "POST", url, r) + if err != nil { + return "", fmt.Errorf("create http request: %w", err) + } - request.Header.Add("Content-Type", "application/json") - request.Header.Add("Accept", "application/json") + request.Header.Add("Content-Type", "application/json") + request.Header.Add("Accept", "application/json") + request.Header.Set("Authorization", "DPoP "+oauthsession.AccessToken) - dpopJwt, err := pdsDpopJwt("POST", url, oauthsession.AuthserverIss, oauthsession.AccessToken, oauthsession.DpopPdsNonce, privateJwk) + privateKey, err := oauthsession.CreatePrivateKey() + if err != nil { + return "", fmt.Errorf("create private key: %w", err) + } + + // try a maximum of 2 times to make the request. If the first attempt fails because the server returns an unauthorized due to a new use_dpop_nonce being issued, + // then try again. Otherwise just try once. + for range 2 { + dpopJwt, err := s.oauthService.PdsDpopJwt("POST", url, oauthsession, privateKey) if err != nil { return "", err } request.Header.Set("DPoP", dpopJwt) - request.Header.Set("Authorization", "DPoP "+oauthsession.AccessToken) resp, err := s.httpClient.Do(request) if err != nil { @@ -76,34 +79,42 @@ func (s *Server) CreateNewStatus(ctx context.Context, oauthsession oauth.Session } defer resp.Body.Close() + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusBadRequest && resp.StatusCode != http.StatusUnauthorized { + return "", fmt.Errorf("unexpected status code returned: %d", resp.StatusCode) + } + + var result CreateRecordResp + err = decodeResp(resp.Body, &result) + if err != nil { + // just log the error. + // if a HTTP 200 is received then the record has been created and we only use the response URI to make an optimistic write to our DB, so nothing will go wrong here. + // if a HTTP 400 then we can at least log return that it was a bad request. + // if a HTTP 401 we only do something if the error string is use_dpop_nonce + slog.Error("decode response body", "error", err) + } + + slog.Info("resp", "status", resp.StatusCode) + if resp.StatusCode == http.StatusOK { - var result CreateRecordResp - err = decodeResp(resp.Body, &result) - if err != nil { - // just log error because we got a 200 indicating that the record was created. If this were to be tried again due to an error - // returned here, there would be duplicate data - slog.Error("decode success response", "error", err) - } return result.URI, nil } - var errorResp XRPCError - err = decodeResp(resp.Body, &errorResp) - if err != nil { - return "", fmt.Errorf("decode error resp: %w", err) + if resp.StatusCode == http.StatusBadRequest { + return "", fmt.Errorf("bad request: %s - %s", result.Message, result.ErrStr) } - if resp.StatusCode == 400 || resp.StatusCode == 401 && errorResp.ErrStr == "use_dpop_nonce" { + if resp.StatusCode == http.StatusUnauthorized && result.ErrStr == "use_dpop_nonce" { newNonce := resp.Header.Get("DPoP-Nonce") oauthsession.DpopPdsNonce = newNonce err := s.oauthService.UpdateOAuthSessionDPopPDSNonce(oauthsession.Did, newNonce) if err != nil { + // just log the error because we can still proceed without storing it. slog.Error("updating oauth session in store with new DPoP PDS nonce", "error", err) } continue } - slog.Error("got error", "status code", resp.StatusCode, "message", errorResp.Message, "error", errorResp.ErrStr) + return "", fmt.Errorf("received an unauthorized status code and message: %s - %s", result.ErrStr, result.Message) } return "", fmt.Errorf("failed to create status record") -- 2.51.2