From d76a775df3685a958b1c5796f3ad013ada3dfb83 Mon Sep 17 00:00:00 2001 From: oppiliappan Date: Sun, 05 Oct 2025 11:33:39 +0000 Subject: [PATCH] appview: switch to indigo oauth library Signed-off-by: oppiliappan --- appview/issues/issues.go | 22 +++++++++++----------- appview/knots/knots.go | 12 ++++++------ appview/labels/labels.go | 18 +++++++++--------- appview/middleware/middleware.go | 19 +++++-------------- appview/notifications/notifications.go | 38 ++++++++++++++++++-------------------- appview/oauth/client/oauth_client.go | 24 ------------------------ appview/oauth/consts.go | 3 ++- appview/oauth/handler.go | 65 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ appview/oauth/handler/handler.go | 538 ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- appview/oauth/oauth.go | 309 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- appview/oauth/store.go | 147 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ appview/pages/templates/layouts/fragments/topbar.html | 2 +- appview/pages/templates/repo/pulls/fragments/pullNewComment.html | 2 +- appview/pages/templates/user/settings/profile.html | 4 +--- appview/pipelines/pipelines.go | 3 ++- appview/pulls/pulls.go | 12 ++++++------ appview/repo/artifact.go | 21 +++++++++++---------- appview/repo/repo.go | 63 ++++++++++++++++++++++++++++----------------------------------- appview/settings/settings.go | 4 ++-- appview/signup/signup.go | 2 -- appview/spindles/spindles.go | 10 +++++----- appview/state/follow.go | 4 ++-- appview/state/login.go | 63 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ appview/state/profile.go | 4 ++-- appview/state/reaction.go | 6 +++--- appview/state/router.go | 15 +++++---------- appview/state/star.go | 4 ++-- appview/state/state.go | 43 +++++++++++++++++++++++-------------------- appview/strings/strings.go | 16 +++++++++------- appview/xrpcclient/xrpc.go | 99 --------------------------------------------------------------------------------------------------- go.mod | 2 +- go.sum | 2 ++ 32 file(s) changed, 539 insertion(s)(+), 1037 deletion(s)(-) diff --git a/appview/issues/issues.go b/appview/issues/issues.go --- a/appview/issues/issues.go +++ b/appview/issues/issues.go @@ -12,6 +12,7 @@ "slices" "time" comatproto "github.com/bluesky-social/indigo/api/atproto" + atpclient "github.com/bluesky-social/indigo/atproto/client" "github.com/bluesky-social/indigo/atproto/syntax" lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/go-chi/chi/v5" @@ -26,7 +27,6 @@ "tangled.org/core/appview/pages" "tangled.org/core/appview/pagination" "tangled.org/core/appview/reporesolver" "tangled.org/core/appview/validator" - "tangled.org/core/appview/xrpcclient" "tangled.org/core/idresolver" tlog "tangled.org/core/log" "tangled.org/core/tid" @@ -166,14 +166,14 @@ rp.pages.Notice(w, noticeId, "Failed to edit issue.") return } - ex, err := client.RepoGetRecord(r.Context(), "", tangled.RepoIssueNSID, user.Did, newIssue.Rkey) + ex, err := comatproto.RepoGetRecord(r.Context(), client, "", tangled.RepoIssueNSID, user.Did, newIssue.Rkey) if err != nil { l.Error("failed to get record", "err", err) rp.pages.Notice(w, noticeId, "Failed to edit issue, no record found on PDS.") return } - _, err = client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + _, err = comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.RepoIssueNSID, Repo: user.Did, Rkey: newIssue.Rkey, @@ -241,7 +241,7 @@ log.Println("failed to get authorized client", err) rp.pages.Notice(w, "issue-comment", "Failed to delete comment.") return } - _, err = client.RepoDeleteRecord(r.Context(), &comatproto.RepoDeleteRecord_Input{ + _, err = comatproto.RepoDeleteRecord(r.Context(), client, &comatproto.RepoDeleteRecord_Input{ Collection: tangled.RepoIssueNSID, Repo: issue.Did, Rkey: issue.Rkey, @@ -408,7 +408,7 @@ return } // create a record first - resp, err := client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + resp, err := comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.RepoIssueCommentNSID, Repo: comment.Did, Rkey: comment.Rkey, @@ -559,14 +559,14 @@ // rkey is optional, it was introduced later if newComment.Rkey != "" { // update the record on pds - ex, err := client.RepoGetRecord(r.Context(), "", tangled.RepoIssueCommentNSID, user.Did, comment.Rkey) + ex, err := comatproto.RepoGetRecord(r.Context(), client, "", tangled.RepoIssueCommentNSID, user.Did, comment.Rkey) if err != nil { log.Println("failed to get record", "err", err, "did", newComment.Did, "rkey", newComment.Rkey) rp.pages.Notice(w, fmt.Sprintf("comment-%s-status", commentId), "Failed to update description, no record found on PDS.") return } - _, err = client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + _, err = comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.RepoIssueCommentNSID, Repo: user.Did, Rkey: newComment.Rkey, @@ -733,7 +733,7 @@ log.Println("failed to get authorized client", err) rp.pages.Notice(w, "issue-comment", "Failed to delete comment.") return } - _, err = client.RepoDeleteRecord(r.Context(), &comatproto.RepoDeleteRecord_Input{ + _, err = comatproto.RepoDeleteRecord(r.Context(), client, &comatproto.RepoDeleteRecord_Input{ Collection: tangled.RepoIssueCommentNSID, Repo: user.Did, Rkey: comment.Rkey, @@ -865,7 +865,7 @@ l.Error("failed to get authorized client", "err", err) rp.pages.Notice(w, "issues", "Failed to create issue.") return } - resp, err := client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + resp, err := comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.RepoIssueNSID, Repo: user.Did, Rkey: issue.Rkey, @@ -923,7 +923,7 @@ // this is used to rollback changes made to the PDS // // it is a no-op if the provided ATURI is empty -func rollbackRecord(ctx context.Context, aturi string, xrpcc *xrpcclient.Client) error { +func rollbackRecord(ctx context.Context, aturi string, client *atpclient.APIClient) error { if aturi == "" { return nil } @@ -934,7 +934,7 @@ collection := parsed.Collection().String() repo := parsed.Authority().String() rkey := parsed.RecordKey().String() - _, err := xrpcc.RepoDeleteRecord(ctx, &comatproto.RepoDeleteRecord_Input{ + _, err := comatproto.RepoDeleteRecord(ctx, client, &comatproto.RepoDeleteRecord_Input{ Collection: collection, Repo: repo, Rkey: rkey, diff --git a/appview/knots/knots.go b/appview/knots/knots.go --- a/appview/knots/knots.go +++ b/appview/knots/knots.go @@ -185,14 +185,14 @@ fail() return } - ex, _ := client.RepoGetRecord(r.Context(), "", tangled.KnotNSID, user.Did, domain) + ex, _ := comatproto.RepoGetRecord(r.Context(), client, "", tangled.KnotNSID, user.Did, domain) var exCid *string if ex != nil { exCid = ex.Cid } // re-announce by registering under same rkey - _, err = client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + _, err = comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.KnotNSID, Repo: user.Did, Rkey: domain, @@ -323,7 +323,7 @@ fail() return } - _, err = client.RepoDeleteRecord(r.Context(), &comatproto.RepoDeleteRecord_Input{ + _, err = comatproto.RepoDeleteRecord(r.Context(), client, &comatproto.RepoDeleteRecord_Input{ Collection: tangled.KnotNSID, Repo: user.Did, Rkey: domain, @@ -431,14 +431,14 @@ fail() return } - ex, _ := client.RepoGetRecord(r.Context(), "", tangled.KnotNSID, user.Did, domain) + ex, _ := comatproto.RepoGetRecord(r.Context(), client, "", tangled.KnotNSID, user.Did, domain) var exCid *string if ex != nil { exCid = ex.Cid } // ignore the error here - _, err = client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + _, err = comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.KnotNSID, Repo: user.Did, Rkey: domain, @@ -555,7 +555,7 @@ } rkey := tid.TID() - _, err = client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + _, err = comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.KnotMemberNSID, Repo: user.Did, Rkey: rkey, diff --git a/appview/labels/labels.go b/appview/labels/labels.go --- a/appview/labels/labels.go +++ b/appview/labels/labels.go @@ -9,11 +9,6 @@ "log/slog" "net/http" "time" - comatproto "github.com/bluesky-social/indigo/api/atproto" - "github.com/bluesky-social/indigo/atproto/syntax" - lexutil "github.com/bluesky-social/indigo/lex/util" - "github.com/go-chi/chi/v5" - "tangled.org/core/api/tangled" "tangled.org/core/appview/db" "tangled.org/core/appview/middleware" @@ -21,10 +16,15 @@ "tangled.org/core/appview/models" "tangled.org/core/appview/oauth" "tangled.org/core/appview/pages" "tangled.org/core/appview/validator" - "tangled.org/core/appview/xrpcclient" "tangled.org/core/log" "tangled.org/core/rbac" "tangled.org/core/tid" + + comatproto "github.com/bluesky-social/indigo/api/atproto" + atpclient "github.com/bluesky-social/indigo/atproto/client" + "github.com/bluesky-social/indigo/atproto/syntax" + lexutil "github.com/bluesky-social/indigo/lex/util" + "github.com/go-chi/chi/v5" ) type Labels struct { @@ -196,7 +196,7 @@ fail("Failed to authorize user.", err) return } - resp, err := client.RepoPutRecord(r.Context(), &comatproto.RepoPutRecord_Input{ + resp, err := comatproto.RepoPutRecord(r.Context(), client, &comatproto.RepoPutRecord_Input{ Collection: tangled.LabelOpNSID, Repo: did, Rkey: rkey, @@ -252,7 +252,7 @@ // this is used to rollback changes made to the PDS // // it is a no-op if the provided ATURI is empty -func rollbackRecord(ctx context.Context, aturi string, xrpcc *xrpcclient.Client) error { +func rollbackRecord(ctx context.Context, aturi string, client *atpclient.APIClient) error { if aturi == "" { return nil } @@ -263,7 +263,7 @@ collection := parsed.Collection().String() repo := parsed.Authority().String() rkey := parsed.RecordKey().String() - _, err := xrpcc.RepoDeleteRecord(ctx, &comatproto.RepoDeleteRecord_Input{ + _, err := comatproto.RepoDeleteRecord(ctx, client, &comatproto.RepoDeleteRecord_Input{ Collection: collection, Repo: repo, Rkey: rkey, diff --git a/appview/middleware/middleware.go b/appview/middleware/middleware.go --- a/appview/middleware/middleware.go +++ b/appview/middleware/middleware.go @@ -43,16 +43,7 @@ } type middlewareFunc func(http.Handler) http.Handler -func (mw *Middleware) TryRefreshSession() middlewareFunc { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _, _ = mw.oauth.GetSession(r) - next.ServeHTTP(w, r) - }) - } -} - -func AuthMiddleware(a *oauth.OAuth) middlewareFunc { +func AuthMiddleware(o *oauth.OAuth) middlewareFunc { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { returnURL := "/" @@ -72,15 +63,15 @@ w.WriteHeader(http.StatusOK) } } - _, auth, err := a.GetSession(r) + sess, err := o.ResumeSession(r) if err != nil { - log.Println("not logged in, redirecting", "err", err) + log.Println("failed to resume session, redirecting...", "err", err, "url", r.URL.String()) redirectFunc(w, r) return } - if !auth { - log.Printf("not logged in, redirecting") + if sess == nil { + log.Printf("session is nil, redirecting...") redirectFunc(w, r) return } diff --git a/appview/notifications/notifications.go b/appview/notifications/notifications.go --- a/appview/notifications/notifications.go +++ b/appview/notifications/notifications.go @@ -1,7 +1,6 @@ package notifications import ( - "fmt" "log" "net/http" "strconv" @@ -31,20 +30,21 @@ func (n *Notifications) Router(mw *middleware.Middleware) http.Handler { r := chi.NewRouter() - r.Use(middleware.AuthMiddleware(n.oauth)) - - r.With(middleware.Paginate).Get("/", n.notificationsPage) - r.Get("/count", n.getUnreadCount) - r.Post("/{id}/read", n.markRead) - r.Post("/read-all", n.markAllRead) - r.Delete("/{id}", n.deleteNotification) + + r.Group(func(r chi.Router) { + r.Use(middleware.AuthMiddleware(n.oauth)) + r.With(middleware.Paginate).Get("/", n.notificationsPage) + r.Post("/{id}/read", n.markRead) + r.Post("/read-all", n.markAllRead) + r.Delete("/{id}", n.deleteNotification) + }) return r } func (n *Notifications) notificationsPage(w http.ResponseWriter, r *http.Request) { - userDid := n.oauth.GetDid(r) + user := n.oauth.GetUser(r) page, ok := r.Context().Value("page").(pagination.Page) if !ok { @@ -54,7 +54,7 @@ } total, err := db.CountNotifications( n.db, - db.FilterEq("recipient_did", userDid), + db.FilterEq("recipient_did", user.Did), ) if err != nil { log.Println("failed to get total notifications:", err) @@ -65,7 +65,7 @@ notifications, err := db.GetNotificationsWithEntities( n.db, page, - db.FilterEq("recipient_did", userDid), + db.FilterEq("recipient_did", user.Did), ) if err != nil { log.Println("failed to get notifications:", err) @@ -73,30 +73,28 @@ n.pages.Error500(w) return } - err = n.db.MarkAllNotificationsRead(r.Context(), userDid) + err = n.db.MarkAllNotificationsRead(r.Context(), user.Did) if err != nil { log.Println("failed to mark notifications as read:", err) } unreadCount := 0 - user := n.oauth.GetUser(r) - if user == nil { - http.Error(w, "Failed to get user", http.StatusInternalServerError) - return - } - - fmt.Println(n.pages.Notifications(w, pages.NotificationsParams{ + n.pages.Notifications(w, pages.NotificationsParams{ LoggedInUser: user, Notifications: notifications, UnreadCount: unreadCount, Page: page, Total: total, - })) + }) } func (n *Notifications) getUnreadCount(w http.ResponseWriter, r *http.Request) { user := n.oauth.GetUser(r) + if user == nil { + return + } + count, err := db.CountNotifications( n.db, db.FilterEq("recipient_did", user.Did), diff --git a/appview/oauth/client/oauth_client.go b/appview/oauth/client/oauth_client.go deleted file mode 100644 --- a/appview/oauth/client/oauth_client.go +++ /dev/null @@ -1,24 +0,0 @@ -package client - -import ( - oauth "tangled.org/anirudh.fi/atproto-oauth" - "tangled.org/anirudh.fi/atproto-oauth/helpers" -) - -type OAuthClient struct { - *oauth.Client -} - -func NewClient(clientId, clientJwk, redirectUri string) (*OAuthClient, error) { - k, err := helpers.ParseJWKFromBytes([]byte(clientJwk)) - if err != nil { - return nil, err - } - - cli, err := oauth.NewClient(oauth.ClientArgs{ - ClientId: clientId, - ClientJwk: k, - RedirectUri: redirectUri, - }) - return &OAuthClient{cli}, err -} diff --git a/appview/oauth/consts.go b/appview/oauth/consts.go --- a/appview/oauth/consts.go +++ b/appview/oauth/consts.go @@ -1,9 +1,10 @@ package oauth const ( - SessionName = "appview-session" + SessionName = "appview-session-v2" SessionHandle = "handle" SessionDid = "did" + SessionId = "id" SessionPds = "pds" SessionAccessJwt = "accessJwt" SessionRefreshJwt = "refreshJwt" diff --git a/appview/oauth/handler.go b/appview/oauth/handler.go new file mode 100644 --- /dev/null +++ b/appview/oauth/handler.go @@ -0,0 +1,65 @@ +package oauth + +import ( + "encoding/json" + "log" + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/lestrrat-go/jwx/v2/jwk" +) + +func (o *OAuth) Router() http.Handler { + r := chi.NewRouter() + + r.Get("/oauth/client-metadata.json", o.clientMetadata) + r.Get("/oauth/jwks.json", o.jwks) + r.Get("/oauth/callback", o.callback) + return r +} + +func (o *OAuth) clientMetadata(w http.ResponseWriter, r *http.Request) { + doc := o.ClientApp.Config.ClientMetadata() + doc.JWKSURI = &o.JwksUri + + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(doc); err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } +} + +func (o *OAuth) jwks(w http.ResponseWriter, r *http.Request) { + jwks := o.Config.OAuth.Jwks + pubKey, err := pubKeyFromJwk(jwks) + if err != nil { + log.Printf("error parsing public key: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + response := map[string]any{ + "keys": []jwk.Key{pubKey}, + } + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + json.NewEncoder(w).Encode(response) +} + +func (o *OAuth) callback(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + sessData, err := o.ClientApp.ProcessCallback(ctx, r.URL.Query()) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + if err := o.SaveSession(w, r, sessData); err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + http.Redirect(w, r, "/", http.StatusFound) +} diff --git a/appview/oauth/handler/handler.go b/appview/oauth/handler/handler.go deleted file mode 100644 --- a/appview/oauth/handler/handler.go +++ /dev/null @@ -1,538 +0,0 @@ -package oauth - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "log" - "net/http" - "net/url" - "slices" - "strings" - "time" - - "github.com/go-chi/chi/v5" - "github.com/gorilla/sessions" - "github.com/lestrrat-go/jwx/v2/jwk" - "github.com/posthog/posthog-go" - "tangled.org/anirudh.fi/atproto-oauth/helpers" - tangled "tangled.org/core/api/tangled" - sessioncache "tangled.org/core/appview/cache/session" - "tangled.org/core/appview/config" - "tangled.org/core/appview/db" - "tangled.org/core/appview/middleware" - "tangled.org/core/appview/oauth" - "tangled.org/core/appview/oauth/client" - "tangled.org/core/appview/pages" - "tangled.org/core/consts" - "tangled.org/core/idresolver" - "tangled.org/core/rbac" - "tangled.org/core/tid" -) - -const ( - oauthScope = "atproto transition:generic" -) - -type OAuthHandler struct { - config *config.Config - pages *pages.Pages - idResolver *idresolver.Resolver - sess *sessioncache.SessionStore - db *db.DB - store *sessions.CookieStore - oauth *oauth.OAuth - enforcer *rbac.Enforcer - posthog posthog.Client -} - -func New( - config *config.Config, - pages *pages.Pages, - idResolver *idresolver.Resolver, - db *db.DB, - sess *sessioncache.SessionStore, - store *sessions.CookieStore, - oauth *oauth.OAuth, - enforcer *rbac.Enforcer, - posthog posthog.Client, -) *OAuthHandler { - return &OAuthHandler{ - config: config, - pages: pages, - idResolver: idResolver, - db: db, - sess: sess, - store: store, - oauth: oauth, - enforcer: enforcer, - posthog: posthog, - } -} - -func (o *OAuthHandler) Router() http.Handler { - r := chi.NewRouter() - - r.Get("/login", o.login) - r.Post("/login", o.login) - - r.With(middleware.AuthMiddleware(o.oauth)).Post("/logout", o.logout) - - r.Get("/oauth/client-metadata.json", o.clientMetadata) - r.Get("/oauth/jwks.json", o.jwks) - r.Get("/oauth/callback", o.callback) - return r -} - -func (o *OAuthHandler) clientMetadata(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - json.NewEncoder(w).Encode(o.oauth.ClientMetadata()) -} - -func (o *OAuthHandler) jwks(w http.ResponseWriter, r *http.Request) { - jwks := o.config.OAuth.Jwks - pubKey, err := pubKeyFromJwk(jwks) - if err != nil { - log.Printf("error parsing public key: %v", err) - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - - response := helpers.CreateJwksResponseObject(pubKey) - - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - json.NewEncoder(w).Encode(response) -} - -func (o *OAuthHandler) login(w http.ResponseWriter, r *http.Request) { - switch r.Method { - case http.MethodGet: - returnURL := r.URL.Query().Get("return_url") - o.pages.Login(w, pages.LoginParams{ - ReturnUrl: returnURL, - }) - case http.MethodPost: - handle := r.FormValue("handle") - - // when users copy their handle from bsky.app, it tends to have these characters around it: - // - // @nelind.dk: - // \u202a ensures that the handle is always rendered left to right and - // \u202c reverts that so the rest of the page renders however it should - handle = strings.TrimPrefix(handle, "\u202a") - handle = strings.TrimSuffix(handle, "\u202c") - - // `@` is harmless - handle = strings.TrimPrefix(handle, "@") - - // basic handle validation - if !strings.Contains(handle, ".") { - log.Println("invalid handle format", "raw", handle) - o.pages.Notice(w, "login-msg", fmt.Sprintf("\"%s\" is an invalid handle. Did you mean %s.bsky.social?", handle, handle)) - return - } - - resolved, err := o.idResolver.ResolveIdent(r.Context(), handle) - if err != nil { - log.Println("failed to resolve handle:", err) - o.pages.Notice(w, "login-msg", fmt.Sprintf("\"%s\" is an invalid handle.", handle)) - return - } - self := o.oauth.ClientMetadata() - oauthClient, err := client.NewClient( - self.ClientID, - o.config.OAuth.Jwks, - self.RedirectURIs[0], - ) - - if err != nil { - log.Println("failed to create oauth client:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - authServer, err := oauthClient.ResolvePdsAuthServer(r.Context(), resolved.PDSEndpoint()) - if err != nil { - log.Println("failed to resolve auth server:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - authMeta, err := oauthClient.FetchAuthServerMetadata(r.Context(), authServer) - if err != nil { - log.Println("failed to fetch auth server metadata:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - dpopKey, err := helpers.GenerateKey(nil) - if err != nil { - log.Println("failed to generate dpop key:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - dpopKeyJson, err := json.Marshal(dpopKey) - if err != nil { - log.Println("failed to marshal dpop key:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - parResp, err := oauthClient.SendParAuthRequest(r.Context(), authServer, authMeta, handle, oauthScope, dpopKey) - if err != nil { - log.Println("failed to send par auth request:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - err = o.sess.SaveRequest(r.Context(), sessioncache.OAuthRequest{ - Did: resolved.DID.String(), - PdsUrl: resolved.PDSEndpoint(), - Handle: handle, - AuthserverIss: authMeta.Issuer, - PkceVerifier: parResp.PkceVerifier, - DpopAuthserverNonce: parResp.DpopAuthserverNonce, - DpopPrivateJwk: string(dpopKeyJson), - State: parResp.State, - ReturnUrl: r.FormValue("return_url"), - }) - if err != nil { - log.Println("failed to save oauth request:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - u, _ := url.Parse(authMeta.AuthorizationEndpoint) - query := url.Values{} - query.Add("client_id", self.ClientID) - query.Add("request_uri", parResp.RequestUri) - u.RawQuery = query.Encode() - o.pages.HxRedirect(w, u.String()) - } -} - -func (o *OAuthHandler) callback(w http.ResponseWriter, r *http.Request) { - state := r.FormValue("state") - - oauthRequest, err := o.sess.GetRequestByState(r.Context(), state) - if err != nil { - log.Println("failed to get oauth request:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - defer func() { - err := o.sess.DeleteRequestByState(r.Context(), state) - if err != nil { - log.Println("failed to delete oauth request for state:", state, err) - } - }() - - error := r.FormValue("error") - errorDescription := r.FormValue("error_description") - if error != "" || errorDescription != "" { - log.Printf("error: %s, %s", error, errorDescription) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - code := r.FormValue("code") - if code == "" { - log.Println("missing code for state: ", state) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - iss := r.FormValue("iss") - if iss == "" { - log.Println("missing iss for state: ", state) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - if iss != oauthRequest.AuthserverIss { - log.Println("mismatched iss:", iss, "!=", oauthRequest.AuthserverIss, "for state:", state) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - self := o.oauth.ClientMetadata() - - oauthClient, err := client.NewClient( - self.ClientID, - o.config.OAuth.Jwks, - self.RedirectURIs[0], - ) - - if err != nil { - log.Println("failed to create oauth client:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - jwk, err := helpers.ParseJWKFromBytes([]byte(oauthRequest.DpopPrivateJwk)) - if err != nil { - log.Println("failed to parse jwk:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - tokenResp, err := oauthClient.InitialTokenRequest( - r.Context(), - code, - oauthRequest.AuthserverIss, - oauthRequest.PkceVerifier, - oauthRequest.DpopAuthserverNonce, - jwk, - ) - if err != nil { - log.Println("failed to get token:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - if tokenResp.Scope != oauthScope { - log.Println("scope doesn't match:", tokenResp.Scope) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - err = o.oauth.SaveSession(w, r, *oauthRequest, tokenResp) - if err != nil { - log.Println("failed to save session:", err) - o.pages.Notice(w, "login-msg", "Failed to authenticate. Try again later.") - return - } - - log.Println("session saved successfully") - go o.addToDefaultKnot(oauthRequest.Did) - go o.addToDefaultSpindle(oauthRequest.Did) - - if !o.config.Core.Dev { - err = o.posthog.Enqueue(posthog.Capture{ - DistinctId: oauthRequest.Did, - Event: "signin", - }) - if err != nil { - log.Println("failed to enqueue posthog event:", err) - } - } - - returnUrl := oauthRequest.ReturnUrl - if returnUrl == "" { - returnUrl = "/" - } - - http.Redirect(w, r, returnUrl, http.StatusFound) -} - -func (o *OAuthHandler) logout(w http.ResponseWriter, r *http.Request) { - err := o.oauth.ClearSession(r, w) - if err != nil { - log.Println("failed to clear session:", err) - http.Redirect(w, r, "/", http.StatusFound) - return - } - - log.Println("session cleared successfully") - o.pages.HxRedirect(w, "/login") -} - -func pubKeyFromJwk(jwks string) (jwk.Key, error) { - k, err := helpers.ParseJWKFromBytes([]byte(jwks)) - if err != nil { - return nil, err - } - pubKey, err := k.PublicKey() - if err != nil { - return nil, err - } - return pubKey, nil -} - -func (o *OAuthHandler) addToDefaultSpindle(did string) { - // use the tangled.sh app password to get an accessJwt - // and create an sh.tangled.spindle.member record with that - spindleMembers, err := db.GetSpindleMembers( - o.db, - db.FilterEq("instance", "spindle.tangled.sh"), - db.FilterEq("subject", did), - ) - if err != nil { - log.Printf("failed to get spindle members for did %s: %v", did, err) - return - } - - if len(spindleMembers) != 0 { - log.Printf("did %s is already a member of the default spindle", did) - return - } - - log.Printf("adding %s to default spindle", did) - session, err := o.createAppPasswordSession(o.config.Core.AppPassword, consts.TangledDid) - if err != nil { - log.Printf("failed to create session: %s", err) - return - } - - record := tangled.SpindleMember{ - LexiconTypeID: "sh.tangled.spindle.member", - Subject: did, - Instance: consts.DefaultSpindle, - CreatedAt: time.Now().Format(time.RFC3339), - } - - if err := session.putRecord(record, tangled.SpindleMemberNSID); err != nil { - log.Printf("failed to add member to default spindle: %s", err) - return - } - - log.Printf("successfully added %s to default spindle", did) -} - -func (o *OAuthHandler) addToDefaultKnot(did string) { - // use the tangled.sh app password to get an accessJwt - // and create an sh.tangled.spindle.member record with that - - allKnots, err := o.enforcer.GetKnotsForUser(did) - if err != nil { - log.Printf("failed to get knot members for did %s: %v", did, err) - return - } - - if slices.Contains(allKnots, consts.DefaultKnot) { - log.Printf("did %s is already a member of the default knot", did) - return - } - - log.Printf("adding %s to default knot", did) - session, err := o.createAppPasswordSession(o.config.Core.TmpAltAppPassword, consts.IcyDid) - if err != nil { - log.Printf("failed to create session: %s", err) - return - } - - record := tangled.KnotMember{ - LexiconTypeID: "sh.tangled.knot.member", - Subject: did, - Domain: consts.DefaultKnot, - CreatedAt: time.Now().Format(time.RFC3339), - } - - if err := session.putRecord(record, tangled.KnotMemberNSID); err != nil { - log.Printf("failed to add member to default knot: %s", err) - return - } - - if err := o.enforcer.AddKnotMember(consts.DefaultKnot, did); err != nil { - log.Printf("failed to set up enforcer rules: %s", err) - return - } - - log.Printf("successfully added %s to default Knot", did) -} - -// create a session using apppasswords -type session struct { - AccessJwt string `json:"accessJwt"` - PdsEndpoint string - Did string -} - -func (o *OAuthHandler) createAppPasswordSession(appPassword, did string) (*session, error) { - if appPassword == "" { - return nil, fmt.Errorf("no app password configured, skipping member addition") - } - - resolved, err := o.idResolver.ResolveIdent(context.Background(), did) - if err != nil { - return nil, fmt.Errorf("failed to resolve tangled.sh DID %s: %v", did, err) - } - - pdsEndpoint := resolved.PDSEndpoint() - if pdsEndpoint == "" { - return nil, fmt.Errorf("no PDS endpoint found for tangled.sh DID %s", did) - } - - sessionPayload := map[string]string{ - "identifier": did, - "password": appPassword, - } - sessionBytes, err := json.Marshal(sessionPayload) - if err != nil { - return nil, fmt.Errorf("failed to marshal session payload: %v", err) - } - - sessionURL := pdsEndpoint + "/xrpc/com.atproto.server.createSession" - sessionReq, err := http.NewRequestWithContext(context.Background(), "POST", sessionURL, bytes.NewBuffer(sessionBytes)) - if err != nil { - return nil, fmt.Errorf("failed to create session request: %v", err) - } - sessionReq.Header.Set("Content-Type", "application/json") - - client := &http.Client{Timeout: 30 * time.Second} - sessionResp, err := client.Do(sessionReq) - if err != nil { - return nil, fmt.Errorf("failed to create session: %v", err) - } - defer sessionResp.Body.Close() - - if sessionResp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("failed to create session: HTTP %d", sessionResp.StatusCode) - } - - var session session - if err := json.NewDecoder(sessionResp.Body).Decode(&session); err != nil { - return nil, fmt.Errorf("failed to decode session response: %v", err) - } - - session.PdsEndpoint = pdsEndpoint - session.Did = did - - return &session, nil -} - -func (s *session) putRecord(record any, collection string) error { - recordBytes, err := json.Marshal(record) - if err != nil { - return fmt.Errorf("failed to marshal knot member record: %w", err) - } - - payload := map[string]any{ - "repo": s.Did, - "collection": collection, - "rkey": tid.TID(), - "record": json.RawMessage(recordBytes), - } - - payloadBytes, err := json.Marshal(payload) - if err != nil { - return fmt.Errorf("failed to marshal request payload: %w", err) - } - - url := s.PdsEndpoint + "/xrpc/com.atproto.repo.putRecord" - req, err := http.NewRequestWithContext(context.Background(), "POST", url, bytes.NewBuffer(payloadBytes)) - if err != nil { - return fmt.Errorf("failed to create HTTP request: %w", err) - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+s.AccessJwt) - - client := &http.Client{Timeout: 30 * time.Second} - resp, err := client.Do(req) - if err != nil { - return fmt.Errorf("failed to add user to default service: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("failed to add user to default service: HTTP %d", resp.StatusCode) - } - - return nil -} diff --git a/appview/oauth/oauth.go b/appview/oauth/oauth.go --- a/appview/oauth/oauth.go +++ b/appview/oauth/oauth.go @@ -1,214 +1,173 @@ package oauth import ( + "errors" "fmt" - "log" "net/http" - "net/url" "time" - indigo_xrpc "github.com/bluesky-social/indigo/xrpc" + comatproto "github.com/bluesky-social/indigo/api/atproto" + "github.com/bluesky-social/indigo/atproto/auth/oauth" + atpclient "github.com/bluesky-social/indigo/atproto/client" + "github.com/bluesky-social/indigo/atproto/syntax" + xrpc "github.com/bluesky-social/indigo/xrpc" "github.com/gorilla/sessions" - oauth "tangled.org/anirudh.fi/atproto-oauth" - "tangled.org/anirudh.fi/atproto-oauth/helpers" - sessioncache "tangled.org/core/appview/cache/session" + "github.com/lestrrat-go/jwx/v2/jwk" "tangled.org/core/appview/config" - "tangled.org/core/appview/oauth/client" - xrpc "tangled.org/core/appview/xrpcclient" ) -type OAuth struct { - store *sessions.CookieStore - config *config.Config - sess *sessioncache.SessionStore -} +func New(config *config.Config) (*OAuth, error) { + + var oauthConfig oauth.ClientConfig + var clientUri string -func NewOAuth(config *config.Config, sess *sessioncache.SessionStore) *OAuth { - return &OAuth{ - store: sessions.NewCookieStore([]byte(config.Core.CookieSecret)), - config: config, - sess: sess, + if config.Core.Dev { + clientUri = "http://127.0.0.1:3000" + callbackUri := clientUri + "/oauth/callback" + oauthConfig = oauth.NewLocalhostConfig(callbackUri, []string{"atproto", "transition:generic"}) + } else { + clientUri = config.Core.AppviewHost + clientId := fmt.Sprintf("%s/oauth/client-metadata.json", clientUri) + callbackUri := clientUri + "/oauth/callback" + oauthConfig = oauth.NewPublicConfig(clientId, callbackUri, []string{"atproto", "transition:generic"}) } + + jwksUri := clientUri + "/oauth/jwks.json" + + authStore, err := NewRedisStore(config.Redis.ToURL()) + if err != nil { + return nil, err + } + + sessStore := sessions.NewCookieStore([]byte(config.Core.CookieSecret)) + + return &OAuth{ + ClientApp: oauth.NewClientApp(&oauthConfig, authStore), + Config: config, + SessStore: sessStore, + JwksUri: jwksUri, + }, nil } -func (o *OAuth) Stores() *sessions.CookieStore { - return o.store +type OAuth struct { + ClientApp *oauth.ClientApp + SessStore *sessions.CookieStore + Config *config.Config + JwksUri string } -func (o *OAuth) SaveSession(w http.ResponseWriter, r *http.Request, oreq sessioncache.OAuthRequest, oresp *oauth.TokenResponse) error { +func (o *OAuth) SaveSession(w http.ResponseWriter, r *http.Request, sessData *oauth.ClientSessionData) error { // first we save the did in the user session - userSession, err := o.store.Get(r, SessionName) + userSession, err := o.SessStore.Get(r, SessionName) if err != nil { return err } - userSession.Values[SessionDid] = oreq.Did - userSession.Values[SessionHandle] = oreq.Handle - userSession.Values[SessionPds] = oreq.PdsUrl + userSession.Values[SessionDid] = sessData.AccountDID.String() + userSession.Values[SessionPds] = sessData.HostURL + userSession.Values[SessionId] = sessData.SessionID userSession.Values[SessionAuthenticated] = true - err = userSession.Save(r, w) + return userSession.Save(r, w) +} + +func (o *OAuth) ResumeSession(r *http.Request) (*oauth.ClientSession, error) { + userSession, err := o.SessStore.Get(r, SessionName) if err != nil { - return fmt.Errorf("error saving user session: %w", err) + return nil, fmt.Errorf("error getting user session: %w", err) } - - // then save the whole thing in the db - session := sessioncache.OAuthSession{ - Did: oreq.Did, - Handle: oreq.Handle, - PdsUrl: oreq.PdsUrl, - DpopAuthserverNonce: oreq.DpopAuthserverNonce, - AuthServerIss: oreq.AuthserverIss, - DpopPrivateJwk: oreq.DpopPrivateJwk, - AccessJwt: oresp.AccessToken, - RefreshJwt: oresp.RefreshToken, - Expiry: time.Now().Add(time.Duration(oresp.ExpiresIn) * time.Second).Format(time.RFC3339), + if userSession.IsNew { + return nil, fmt.Errorf("no session available for user") } - return o.sess.SaveSession(r.Context(), session) -} - -func (o *OAuth) ClearSession(r *http.Request, w http.ResponseWriter) error { - userSession, err := o.store.Get(r, SessionName) - if err != nil || userSession.IsNew { - return fmt.Errorf("error getting user session (or new session?): %w", err) + d := userSession.Values[SessionDid].(string) + sessDid, err := syntax.ParseDID(d) + if err != nil { + return nil, fmt.Errorf("malformed DID in session cookie '%s': %w", d, err) } - did := userSession.Values[SessionDid].(string) + sessId := userSession.Values[SessionId].(string) - err = o.sess.DeleteSession(r.Context(), did) + clientSess, err := o.ClientApp.ResumeSession(r.Context(), sessDid, sessId) if err != nil { - return fmt.Errorf("error deleting oauth session: %w", err) + return nil, fmt.Errorf("failed to resume session: %w", err) } - userSession.Options.MaxAge = -1 - - return userSession.Save(r, w) + return clientSess, nil } -func (o *OAuth) GetSession(r *http.Request) (*sessioncache.OAuthSession, bool, error) { - userSession, err := o.store.Get(r, SessionName) - if err != nil || userSession.IsNew { - return nil, false, fmt.Errorf("error getting user session (or new session?): %w", err) +func (o *OAuth) DeleteSession(w http.ResponseWriter, r *http.Request) error { + userSession, err := o.SessStore.Get(r, SessionName) + if err != nil { + return fmt.Errorf("error getting user session: %w", err) } - - did := userSession.Values[SessionDid].(string) - auth := userSession.Values[SessionAuthenticated].(bool) - - session, err := o.sess.GetSession(r.Context(), did) - if err != nil { - return nil, false, fmt.Errorf("error getting oauth session: %w", err) + if userSession.IsNew { + return fmt.Errorf("no session available for user") } - expiry, err := time.Parse(time.RFC3339, session.Expiry) + d := userSession.Values[SessionDid].(string) + sessDid, err := syntax.ParseDID(d) if err != nil { - return nil, false, fmt.Errorf("error parsing expiry time: %w", err) + return fmt.Errorf("malformed DID in session cookie '%s': %w", d, err) } - if time.Until(expiry) <= 5*time.Minute { - privateJwk, err := helpers.ParseJWKFromBytes([]byte(session.DpopPrivateJwk)) - if err != nil { - return nil, false, err - } - self := o.ClientMetadata() + sessId := userSession.Values[SessionId].(string) - oauthClient, err := client.NewClient( - self.ClientID, - o.config.OAuth.Jwks, - self.RedirectURIs[0], - ) + // delete the session + err1 := o.ClientApp.Logout(r.Context(), sessDid, sessId) - if err != nil { - return nil, false, err - } + // remove the cookie + userSession.Options.MaxAge = -1 + err2 := o.SessStore.Save(r, w, userSession) - resp, err := oauthClient.RefreshTokenRequest(r.Context(), session.RefreshJwt, session.AuthServerIss, session.DpopAuthserverNonce, privateJwk) - if err != nil { - return nil, false, err - } + return errors.Join(err1, err2) +} - newExpiry := time.Now().Add(time.Duration(resp.ExpiresIn) * time.Second).Format(time.RFC3339) - err = o.sess.RefreshSession(r.Context(), did, resp.AccessToken, resp.RefreshToken, newExpiry) - if err != nil { - return nil, false, fmt.Errorf("error refreshing oauth session: %w", err) - } - - // update the current session - session.AccessJwt = resp.AccessToken - session.RefreshJwt = resp.RefreshToken - session.DpopAuthserverNonce = resp.DpopAuthserverNonce - session.Expiry = newExpiry +func pubKeyFromJwk(jwks string) (jwk.Key, error) { + k, err := jwk.ParseKey([]byte(jwks)) + if err != nil { + return nil, err + } + pubKey, err := k.PublicKey() + if err != nil { + return nil, err } - - return session, auth, nil + return pubKey, nil } type User struct { - Handle string - Did string - Pds string + Did string + Pds string } -func (a *OAuth) GetUser(r *http.Request) *User { - clientSession, err := a.store.Get(r, SessionName) +func (o *OAuth) GetUser(r *http.Request) *User { + sess, err := o.SessStore.Get(r, SessionName) - if err != nil || clientSession.IsNew { + if err != nil || sess.IsNew { return nil } return &User{ - Handle: clientSession.Values[SessionHandle].(string), - Did: clientSession.Values[SessionDid].(string), - Pds: clientSession.Values[SessionPds].(string), + Did: sess.Values[SessionDid].(string), + Pds: sess.Values[SessionPds].(string), } } -func (a *OAuth) GetDid(r *http.Request) string { - clientSession, err := a.store.Get(r, SessionName) - - if err != nil || clientSession.IsNew { - return "" +func (o *OAuth) GetDid(r *http.Request) string { + if u := o.GetUser(r); u != nil { + return u.Did } - return clientSession.Values[SessionDid].(string) + return "" } -func (o *OAuth) AuthorizedClient(r *http.Request) (*xrpc.Client, error) { - session, auth, err := o.GetSession(r) +func (o *OAuth) AuthorizedClient(r *http.Request) (*atpclient.APIClient, error) { + session, err := o.ResumeSession(r) if err != nil { return nil, fmt.Errorf("error getting session: %w", err) } - if !auth { - return nil, fmt.Errorf("not authorized") - } - - client := &oauth.XrpcClient{ - OnDpopPdsNonceChanged: func(did, newNonce string) { - err := o.sess.UpdateNonce(r.Context(), did, newNonce) - if err != nil { - log.Printf("error updating dpop pds nonce: %v", err) - } - }, - } - - privateJwk, err := helpers.ParseJWKFromBytes([]byte(session.DpopPrivateJwk)) - if err != nil { - return nil, fmt.Errorf("error parsing private jwk: %w", err) - } - - xrpcClient := xrpc.NewClient(client, &oauth.XrpcAuthedRequestArgs{ - Did: session.Did, - PdsUrl: session.PdsUrl, - DpopPdsNonce: session.PdsUrl, - AccessToken: session.AccessJwt, - Issuer: session.AuthServerIss, - DpopPrivateJwk: privateJwk, - }) - - return xrpcClient, nil + return session.APIClient(), nil } -// use this to create a client to communicate with knots or spindles -// // this is a higher level abstraction on ServerGetServiceAuth type ServiceClientOpts struct { service string @@ -259,13 +218,13 @@ return scheme + s.service } -func (o *OAuth) ServiceClient(r *http.Request, os ...ServiceClientOpt) (*indigo_xrpc.Client, error) { +func (o *OAuth) ServiceClient(r *http.Request, os ...ServiceClientOpt) (*xrpc.Client, error) { opts := ServiceClientOpts{} for _, o := range os { o(&opts) } - authorizedClient, err := o.AuthorizedClient(r) + client, err := o.AuthorizedClient(r) if err != nil { return nil, err } @@ -276,13 +235,13 @@ if opts.exp < sixty { opts.exp = sixty } - resp, err := authorizedClient.ServerGetServiceAuth(r.Context(), opts.Audience(), opts.exp, opts.lxm) + resp, err := comatproto.ServerGetServiceAuth(r.Context(), client, opts.Audience(), opts.exp, opts.lxm) if err != nil { return nil, err } - return &indigo_xrpc.Client{ - Auth: &indigo_xrpc.AuthInfo{ + return &xrpc.Client{ + Auth: &xrpc.AuthInfo{ AccessJwt: resp.Token, }, Host: opts.Host(), @@ -291,57 +250,3 @@ Timeout: time.Second * 5, }, }, nil } - -type ClientMetadata struct { - ClientID string `json:"client_id"` - ClientName string `json:"client_name"` - SubjectType string `json:"subject_type"` - ClientURI string `json:"client_uri"` - RedirectURIs []string `json:"redirect_uris"` - GrantTypes []string `json:"grant_types"` - ResponseTypes []string `json:"response_types"` - ApplicationType string `json:"application_type"` - DpopBoundAccessTokens bool `json:"dpop_bound_access_tokens"` - JwksURI string `json:"jwks_uri"` - Scope string `json:"scope"` - TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"` - TokenEndpointAuthSigningAlg string `json:"token_endpoint_auth_signing_alg"` -} - -func (o *OAuth) ClientMetadata() ClientMetadata { - makeRedirectURIs := func(c string) []string { - return []string{fmt.Sprintf("%s/oauth/callback", c)} - } - - clientURI := o.config.Core.AppviewHost - clientID := fmt.Sprintf("%s/oauth/client-metadata.json", clientURI) - redirectURIs := makeRedirectURIs(clientURI) - - if o.config.Core.Dev { - clientURI = "http://127.0.0.1:3000" - redirectURIs = makeRedirectURIs(clientURI) - - query := url.Values{} - query.Add("redirect_uri", redirectURIs[0]) - query.Add("scope", "atproto transition:generic") - clientID = fmt.Sprintf("http://localhost?%s", query.Encode()) - } - - jwksURI := fmt.Sprintf("%s/oauth/jwks.json", clientURI) - - return ClientMetadata{ - ClientID: clientID, - ClientName: "Tangled", - SubjectType: "public", - ClientURI: clientURI, - RedirectURIs: redirectURIs, - GrantTypes: []string{"authorization_code", "refresh_token"}, - ResponseTypes: []string{"code"}, - ApplicationType: "web", - DpopBoundAccessTokens: true, - JwksURI: jwksURI, - Scope: "atproto transition:generic", - TokenEndpointAuthMethod: "private_key_jwt", - TokenEndpointAuthSigningAlg: "ES256", - } -} diff --git a/appview/oauth/store.go b/appview/oauth/store.go new file mode 100644 --- /dev/null +++ b/appview/oauth/store.go @@ -0,0 +1,147 @@ +package oauth + +import ( + "context" + "encoding/json" + "fmt" + "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/redis/go-redis/v9" +) + +// redis-backed implementation of ClientAuthStore. +type RedisStore struct { + client *redis.Client + SessionTTL time.Duration + AuthRequestTTL time.Duration +} + +var _ oauth.ClientAuthStore = &RedisStore{} + +func NewRedisStore(redisURL string) (*RedisStore, error) { + opts, err := redis.ParseURL(redisURL) + if err != nil { + return nil, fmt.Errorf("failed to parse redis URL: %w", err) + } + + client := redis.NewClient(opts) + + // test the connection + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + if err := client.Ping(ctx).Err(); err != nil { + return nil, fmt.Errorf("failed to connect to redis: %w", err) + } + + return &RedisStore{ + client: client, + SessionTTL: 30 * 24 * time.Hour, // 30 days + AuthRequestTTL: 10 * time.Minute, // 10 minutes + }, nil +} + +func (r *RedisStore) Close() error { + return r.client.Close() +} + +func sessionKey(did syntax.DID, sessionID string) string { + return fmt.Sprintf("oauth:session:%s:%s", did, sessionID) +} + +func authRequestKey(state string) string { + return fmt.Sprintf("oauth:auth_request:%s", state) +} + +func (r *RedisStore) GetSession(ctx context.Context, did syntax.DID, sessionID string) (*oauth.ClientSessionData, error) { + key := sessionKey(did, sessionID) + data, err := r.client.Get(ctx, key).Bytes() + if err == redis.Nil { + return nil, fmt.Errorf("session not found: %s", did) + } + if err != nil { + return nil, fmt.Errorf("failed to get session: %w", err) + } + + var sess oauth.ClientSessionData + if err := json.Unmarshal(data, &sess); err != nil { + return nil, fmt.Errorf("failed to unmarshal session: %w", err) + } + + return &sess, nil +} + +func (r *RedisStore) SaveSession(ctx context.Context, sess oauth.ClientSessionData) error { + key := sessionKey(sess.AccountDID, sess.SessionID) + + data, err := json.Marshal(sess) + if err != nil { + return fmt.Errorf("failed to marshal session: %w", err) + } + + if err := r.client.Set(ctx, key, data, r.SessionTTL).Err(); err != nil { + return fmt.Errorf("failed to save session: %w", err) + } + + return nil +} + +func (r *RedisStore) DeleteSession(ctx context.Context, did syntax.DID, sessionID string) error { + key := sessionKey(did, sessionID) + if err := r.client.Del(ctx, key).Err(); err != nil { + return fmt.Errorf("failed to delete session: %w", err) + } + return nil +} + +func (r *RedisStore) GetAuthRequestInfo(ctx context.Context, state string) (*oauth.AuthRequestData, error) { + key := authRequestKey(state) + data, err := r.client.Get(ctx, key).Bytes() + if err == redis.Nil { + return nil, fmt.Errorf("request info not found: %s", state) + } + if err != nil { + return nil, fmt.Errorf("failed to get auth request: %w", err) + } + + var req oauth.AuthRequestData + if err := json.Unmarshal(data, &req); err != nil { + return nil, fmt.Errorf("failed to unmarshal auth request: %w", err) + } + + return &req, nil +} + +func (r *RedisStore) SaveAuthRequestInfo(ctx context.Context, info oauth.AuthRequestData) error { + key := authRequestKey(info.State) + + // check if already exists (to match MemStore behavior) + exists, err := r.client.Exists(ctx, key).Result() + if err != nil { + return fmt.Errorf("failed to check auth request existence: %w", err) + } + if exists > 0 { + return fmt.Errorf("auth request already saved for state %s", info.State) + } + + data, err := json.Marshal(info) + if err != nil { + return fmt.Errorf("failed to marshal auth request: %w", err) + } + + if err := r.client.Set(ctx, key, data, r.AuthRequestTTL).Err(); err != nil { + return fmt.Errorf("failed to save auth request: %w", err) + } + + return nil +} + +func (r *RedisStore) DeleteAuthRequestInfo(ctx context.Context, state string) error { + key := authRequestKey(state) + if err := r.client.Del(ctx, key).Err(); err != nil { + return fmt.Errorf("failed to delete auth request: %w", err) + } + return nil +} diff --git a/appview/pages/templates/layouts/fragments/topbar.html b/appview/pages/templates/layouts/fragments/topbar.html --- a/appview/pages/templates/layouts/fragments/topbar.html +++ b/appview/pages/templates/layouts/fragments/topbar.html @@ -51,7 +51,7 @@