From ce0709d273d7daa44889974c5ec71835aa43131d Mon Sep 17 00:00:00 2001 From: Anirudh Oppiliappan Date: Tue, 28 Jan 2025 12:26:52 +0200 Subject: [PATCH] clean up auth code --- legit/routes/auth.go | 23 +++------ legit/routes/auth/auth.go | 98 ++++++++++++++++++++++++++++++++++++++ legit/routes/auth/types.go | 15 ++++++ legit/routes/handler.go | 12 +++-- legit/routes/routes.go | 62 ++++++------------------ 5 files changed, 142 insertions(+), 68 deletions(-) create mode 100644 legit/routes/auth/auth.go create mode 100644 legit/routes/auth/types.go diff --git a/legit/routes/auth.go b/legit/routes/auth.go index ed94cf2..b3089eb 100644 --- a/legit/routes/auth.go +++ b/legit/routes/auth.go @@ -7,6 +7,7 @@ import ( comatproto "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/xrpc" + rauth "github.com/icyphox/bild/legit/routes/auth" ) const ( @@ -19,7 +20,7 @@ func (h *Handle) AuthMiddleware(next http.Handler) http.Handler { auth, ok := session.Values["authenticated"].(bool) if !ok || !auth { - http.Error(w, "Forbidden: You are not logged in", http.StatusForbidden) + http.Redirect(w, r, "/login", http.StatusTemporaryRedirect) return } @@ -43,28 +44,16 @@ func (h *Handle) AuthMiddleware(next http.Handler) http.Handler { }, } atSession, err := comatproto.ServerRefreshSession(r.Context(), &client) - if err != nil { log.Println(err) - http.Error(w, "Internal Server Error", http.StatusInternalServerError) + h.Write500(w) return } - clientSession, _ := h.s.Get(r, "bild-session") - clientSession.Values["handle"] = atSession.Handle - clientSession.Values["did"] = atSession.Did - clientSession.Values["accessJwt"] = atSession.AccessJwt - clientSession.Values["refreshJwt"] = atSession.RefreshJwt - clientSession.Values["expiry"] = time.Now().Add(time.Hour).String() - clientSession.Values["pds"] = pdsUrl - clientSession.Values["authenticated"] = true - - err = clientSession.Save(r, w) - + err = h.auth.StoreSession(r, w, nil, &rauth.AtSessionRefresh{ServerRefreshSession_Output: *atSession, PDSEndpoint: pdsUrl}) if err != nil { - log.Printf("failed to store session for did: %s\n", atSession.Did) - log.Println(err) - http.Error(w, "Internal Server Error", http.StatusInternalServerError) + log.Printf("failed to store session for did: %s\n: %s", atSession.Did, err) + h.Write500(w) return } diff --git a/legit/routes/auth/auth.go b/legit/routes/auth/auth.go new file mode 100644 index 0000000..b33ca40 --- /dev/null +++ b/legit/routes/auth/auth.go @@ -0,0 +1,98 @@ +package auth + +import ( + "context" + "fmt" + "net/http" + "time" + + comatproto "github.com/bluesky-social/indigo/api/atproto" + "github.com/bluesky-social/indigo/atproto/identity" + "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/bluesky-social/indigo/xrpc" + "github.com/gorilla/sessions" +) + +type Auth struct { + s sessions.Store +} + +func NewAuth(store sessions.Store) *Auth { + return &Auth{store} +} + +func resolveIdent(ctx context.Context, arg string) (*identity.Identity, error) { + id, err := syntax.ParseAtIdentifier(arg) + if err != nil { + return nil, err + } + + dir := identity.DefaultDirectory() + return dir.Lookup(ctx, *id) +} + +func (a *Auth) CreateInitialSession(w http.ResponseWriter, r *http.Request, username, appPassword string) (AtSessionCreate, error) { + ctx := r.Context() + resolved, err := resolveIdent(ctx, username) + if err != nil { + return AtSessionCreate{}, fmt.Errorf("invalid handle: %s", err) + } + + pdsUrl := resolved.PDSEndpoint() + client := xrpc.Client{ + Host: pdsUrl, + } + + atSession, err := comatproto.ServerCreateSession(ctx, &client, &comatproto.ServerCreateSession_Input{ + Identifier: resolved.DID.String(), + Password: appPassword, + }) + if err != nil { + return AtSessionCreate{}, fmt.Errorf("invalid app password") + } + + return AtSessionCreate{ + ServerCreateSession_Output: *atSession, + PDSEndpoint: pdsUrl, + }, nil +} + +func (a *Auth) StoreSession(r *http.Request, w http.ResponseWriter, atSessionCreate *AtSessionCreate, atSessionRefresh *AtSessionRefresh) error { + if atSessionCreate != nil { + atSession := atSessionCreate + + clientSession, _ := a.s.Get(r, "bild-session") + clientSession.Values["handle"] = atSession.Handle + clientSession.Values["did"] = atSession.Did + clientSession.Values["accessJwt"] = atSession.AccessJwt + clientSession.Values["refreshJwt"] = atSession.RefreshJwt + clientSession.Values["expiry"] = time.Now().Add(time.Hour).String() + clientSession.Values["pds"] = atSession.PDSEndpoint + clientSession.Values["authenticated"] = true + + return clientSession.Save(r, w) + } else { + atSession := atSessionRefresh + + clientSession, _ := a.s.Get(r, "bild-session") + clientSession.Values["handle"] = atSession.Handle + clientSession.Values["did"] = atSession.Did + clientSession.Values["accessJwt"] = atSession.AccessJwt + clientSession.Values["refreshJwt"] = atSession.RefreshJwt + clientSession.Values["expiry"] = time.Now().Add(time.Hour).String() + clientSession.Values["pds"] = atSession.PDSEndpoint + clientSession.Values["authenticated"] = true + + return clientSession.Save(r, w) + } +} + +func (a *Auth) GetSessionUser(r *http.Request) (*identity.Identity, error) { + session, _ := a.s.Get(r, "bild-session") + did, ok := session.Values["did"].(string) + if !ok { + return nil, fmt.Errorf("user is not authenticated") + } + + return resolveIdent(r.Context(), did) +} diff --git a/legit/routes/auth/types.go b/legit/routes/auth/types.go new file mode 100644 index 0000000..ea52a14 --- /dev/null +++ b/legit/routes/auth/types.go @@ -0,0 +1,15 @@ +package auth + +import ( + comatproto "github.com/bluesky-social/indigo/api/atproto" +) + +type AtSessionCreate struct { + comatproto.ServerCreateSession_Output + PDSEndpoint string +} + +type AtSessionRefresh struct { + comatproto.ServerRefreshSession_Output + PDSEndpoint string +} diff --git a/legit/routes/handler.go b/legit/routes/handler.go index d290984..ca60df6 100644 --- a/legit/routes/handler.go +++ b/legit/routes/handler.go @@ -11,6 +11,7 @@ import ( "github.com/gorilla/sessions" "github.com/icyphox/bild/legit/config" "github.com/icyphox/bild/legit/db" + "github.com/icyphox/bild/legit/routes/auth" "github.com/icyphox/bild/legit/routes/tmpl" ) @@ -44,16 +45,19 @@ func Setup(c *config.Config) (http.Handler, error) { return nil, fmt.Errorf("failed to load templates: %w", err) } + auth := auth.NewAuth(s) + db, err := db.Setup(c.Server.DBPath) if err != nil { return nil, fmt.Errorf("failed to setup db: %w", err) } h := Handle{ - c: c, - t: t, - s: s, - db: db, + c: c, + t: t, + s: s, + db: db, + auth: auth, } r.Get("/login", h.Login) diff --git a/legit/routes/routes.go b/legit/routes/routes.go index 8ce2ebd..5aa2870 100644 --- a/legit/routes/routes.go +++ b/legit/routes/routes.go @@ -2,7 +2,6 @@ package routes import ( "compress/gzip" - "context" "errors" "fmt" "html/template" @@ -15,10 +14,6 @@ import ( "strings" "time" - comatproto "github.com/bluesky-social/indigo/api/atproto" - "github.com/bluesky-social/indigo/atproto/identity" - "github.com/bluesky-social/indigo/atproto/syntax" - "github.com/bluesky-social/indigo/xrpc" "github.com/dustin/go-humanize" "github.com/go-chi/chi/v5" "github.com/go-git/go-git/v5/plumbing" @@ -26,15 +21,17 @@ import ( "github.com/icyphox/bild/legit/config" "github.com/icyphox/bild/legit/db" "github.com/icyphox/bild/legit/git" + "github.com/icyphox/bild/legit/routes/auth" "github.com/russross/blackfriday/v2" "golang.org/x/crypto/ssh" ) type Handle struct { - c *config.Config - t *template.Template - s *sessions.CookieStore - db *db.DB + c *config.Config + t *template.Template + s *sessions.CookieStore + db *db.DB + auth *auth.Auth } func (h *Handle) Index(w http.ResponseWriter, r *http.Request) { @@ -440,16 +437,6 @@ func (h *Handle) ServeStatic(w http.ResponseWriter, r *http.Request) { http.ServeFile(w, r, f) } -func resolveIdent(ctx context.Context, arg string) (*identity.Identity, error) { - id, err := syntax.ParseAtIdentifier(arg) - if err != nil { - return nil, err - } - - dir := identity.DefaultDirectory() - return dir.Lookup(ctx, *id) -} - func (h *Handle) Login(w http.ResponseWriter, r *http.Request) { switch r.Method { case http.MethodGet: @@ -458,45 +445,26 @@ func (h *Handle) Login(w http.ResponseWriter, r *http.Request) { return } case http.MethodPost: - ctx := r.Context() username := r.FormValue("username") appPassword := r.FormValue("app_password") - resolved, err := resolveIdent(ctx, username) + atSession, err := h.auth.CreateInitialSession(w, r, username, appPassword) if err != nil { - http.Error(w, "invalid `handle`", http.StatusBadRequest) + h.WriteOOBNotice(w, "login", "Invalid username or app password.") + log.Printf("creating initial session: %s", err) return } - pdsUrl := resolved.PDSEndpoint() - client := xrpc.Client{ - Host: pdsUrl, - } - - atSession, err := comatproto.ServerCreateSession(ctx, &client, &comatproto.ServerCreateSession_Input{ - Identifier: resolved.DID.String(), - Password: appPassword, - }) - - clientSession, _ := h.s.Get(r, "bild-session") - clientSession.Values["handle"] = atSession.Handle - clientSession.Values["did"] = atSession.Did - clientSession.Values["accessJwt"] = atSession.AccessJwt - clientSession.Values["refreshJwt"] = atSession.RefreshJwt - clientSession.Values["expiry"] = time.Now().Add(time.Hour).String() - clientSession.Values["pds"] = pdsUrl - clientSession.Values["authenticated"] = true - - err = clientSession.Save(r, w) - + err = h.auth.StoreSession(r, w, &atSession, nil) if err != nil { - log.Printf("failed to store session for did: %s\n", atSession.Did) - log.Println(err) + h.WriteOOBNotice(w, "login", "Failed to store session.") + log.Printf("storing session: %s", err) return } log.Printf("successfully saved session for %s (%s)", atSession.Handle, atSession.Did) - http.Redirect(w, r, "/@"+atSession.Handle, 302) + w.Header().Set("HX-Redirect", "/") + w.WriteHeader(http.StatusOK) } } @@ -508,8 +476,8 @@ func (h *Handle) Keys(w http.ResponseWriter, r *http.Request) { case http.MethodGet: keys, err := h.db.GetPublicKeys(did) if err != nil { + h.WriteOOBNotice(w, "keys", "Failed to list keys. Try again later.") log.Println(err) - http.Error(w, "invalid `did`", http.StatusBadRequest) return } -- 2.51.2