From 4bd740d8a039fa15429f255f33963d571a8ba4a2 Mon Sep 17 00:00:00 2001 From: Luna Seemann Date: Tue, 19 May 2026 21:39:59 +0200 Subject: [PATCH] feat: account switcher (#79) * feat: account switcher * fix(account-switch): merge redirect query params safely Parse redirect targets and query_params with net/url, then merge into a single encoded query string to avoid malformed URLs when next already has a query or query_params starts with ?. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * perf(session): avoid duplicate account lookups Reuse a single session-account fetch path for signin/account/oauth authorize flows by returning both the active repo and account list from one helper. This removes repeated per-account queries on page render while preserving existing behavior. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix(auth): distinguish unauthenticated vs backend session errors Introduce ErrSessionUnauthenticated and treat only that case as a signin redirect. Return server errors for account/session lookup failures in account and oauth authorize/revoke flows so backend issues are not masked as re-login prompts. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix(pr-review): address remaining account/oath review issues Populate authorize/account template render data for all paths, harden account switch against cross-site POSTs, and apply consistent account session cookie options on save. Also fix pointer-to-range-variable in session account lookup. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix(pr-review): resolve remaining template/session threads Use explicit .Repo.Did in account switcher templates to avoid ambiguous embedded Did fields in RepoActor. Reuse the already-loaded session in oauth authorize by adding a helper variant that accepts an existing session instead of re-fetching it. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- server/account_sessions.go | 168 +++++++++++++++++++++++++++++++ server/handle_account.go | 21 +++- server/handle_account_revoke.go | 5 + server/handle_account_signin.go | 72 +++++++++---- server/handle_account_signout.go | 15 +-- server/handle_account_switch.go | 116 +++++++++++++++++++++ server/handle_oauth_authorize.go | 41 +++++++- server/server.go | 1 + server/session_options.go | 16 +++ server/templates/account.html | 17 +++- server/templates/authorize.html | 20 +++- server/templates/signin.html | 3 + test.go | 2 +- 13 files changed, 460 insertions(+), 37 deletions(-) create mode 100644 server/account_sessions.go create mode 100644 server/handle_account_switch.go create mode 100644 server/session_options.go diff --git a/server/account_sessions.go b/server/account_sessions.go new file mode 100644 index 0000000..4fc9b4d --- /dev/null +++ b/server/account_sessions.go @@ -0,0 +1,168 @@ +package server + +import ( + "context" + "errors" + "slices" + "strings" + + "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/gorilla/sessions" + "github.com/haileyok/cocoon/models" + "gorm.io/gorm" +) + +const ( + sessionDidKey = "did" + sessionDidsKey = "dids" +) + +func normalizeSessionDids(dids []string) []string { + normalized := make([]string, 0, len(dids)) + for _, did := range dids { + if did == "" || slices.Contains(normalized, did) { + continue + } + normalized = append(normalized, did) + } + return normalized +} + +func getSessionDids(sess *sessions.Session) []string { + if sess == nil { + return nil + } + + if val, ok := sess.Values[sessionDidsKey]; ok { + switch dids := val.(type) { + case []string: + return normalizeSessionDids(dids) + case []any: + out := make([]string, 0, len(dids)) + for _, did := range dids { + if s, ok := did.(string); ok { + out = append(out, s) + } + } + return normalizeSessionDids(out) + } + } + + if did, ok := sess.Values[sessionDidKey].(string); ok && did != "" { + return []string{did} + } + + return nil +} + +func setSessionDids(sess *sessions.Session, dids []string) { + if sess == nil { + return + } + + normalized := normalizeSessionDids(dids) + if len(normalized) == 0 { + delete(sess.Values, sessionDidKey) + delete(sess.Values, sessionDidsKey) + return + } + + sess.Values[sessionDidsKey] = normalized + if activeDid, ok := sess.Values[sessionDidKey].(string); !ok || !slices.Contains(normalized, activeDid) { + sess.Values[sessionDidKey] = normalized[0] + } +} + +func getActiveSessionDid(sess *sessions.Session) string { + if sess == nil { + return "" + } + + dids := getSessionDids(sess) + if len(dids) == 0 { + return "" + } + + if activeDid, ok := sess.Values[sessionDidKey].(string); ok && slices.Contains(dids, activeDid) { + return activeDid + } + return dids[0] +} + +func setActiveSessionDid(sess *sessions.Session, did string) bool { + if sess == nil || did == "" { + return false + } + + dids := getSessionDids(sess) + if !slices.Contains(dids, did) { + dids = append(dids, did) + } + setSessionDids(sess, dids) + + current, _ := sess.Values[sessionDidKey].(string) + if current == did { + return false + } + sess.Values[sessionDidKey] = did + return true +} + +func removeSessionDid(sess *sessions.Session, did string) { + if sess == nil || did == "" { + return + } + + next := make([]string, 0) + for _, existingDid := range getSessionDids(sess) { + if existingDid != did { + next = append(next, existingDid) + } + } + setSessionDids(sess, next) +} + +func (s *Server) getSessionAccountActors(ctx context.Context, sess *sessions.Session) ([]models.RepoActor, bool, error) { + changed := false + validDids := make([]string, 0) + var accounts []models.RepoActor + for _, did := range getSessionDids(sess) { + repo, err := s.getRepoActorByDid(ctx, did) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + changed = true + continue + } + return nil, changed, err + } + validDids = append(validDids, did) + accounts = append(accounts, *repo) + } + + if changed { + setSessionDids(sess, validDids) + } + return accounts, changed, nil +} + +func (s *Server) resolveLoginHintToDid(ctx context.Context, loginHint string) (string, error) { + loginHint = strings.TrimSpace(loginHint) + if loginHint == "" { + return "", gorm.ErrRecordNotFound + } + + if _, err := syntax.ParseDID(loginHint); err == nil { + return loginHint, nil + } + + normalizedHandle := strings.ToLower(loginHint) + if _, err := syntax.ParseHandle(normalizedHandle); err == nil { + actor, err := s.getActorByHandle(ctx, normalizedHandle) + if err != nil { + return "", err + } + return actor.Did, nil + } + + return "", gorm.ErrRecordNotFound +} diff --git a/server/handle_account.go b/server/handle_account.go index 8ab4450..8a26ef4 100644 --- a/server/handle_account.go +++ b/server/handle_account.go @@ -1,8 +1,10 @@ package server import ( + "errors" "time" + "github.com/haileyok/cocoon/internal/helpers" "github.com/haileyok/cocoon/oauth" "github.com/haileyok/cocoon/oauth/constants" "github.com/haileyok/cocoon/oauth/provider" @@ -14,8 +16,11 @@ func (s *Server) handleAccount(e echo.Context) error { ctx := e.Request().Context() logger := s.logger.With("name", "handleAuth") - repo, sess, err := s.getSessionRepoOrErr(e) + repo, sess, accounts, err := s.getSessionRepoAndAccountsOrErr(e) if err != nil { + if !errors.Is(err, ErrSessionUnauthenticated) { + return helpers.ServerError(e, nil) + } return e.Redirect(303, "/account/signin") } @@ -27,7 +32,11 @@ func (s *Server) handleAccount(e echo.Context) error { sess.AddFlash("Unable to fetch sessions. See server logs for more details.", "error") sess.Save(e.Request(), e.Response()) return e.Render(200, "account.html", map[string]any{ - "flashes": getFlashesFromSession(e, sess), + "Repo": repo, + "Tokens": []map[string]string{}, + "flashes": getFlashesFromSession(e, sess), + "Accounts": accounts, + "ActiveDid": repo.Repo.Did, }) } @@ -69,8 +78,10 @@ func (s *Server) handleAccount(e echo.Context) error { } return e.Render(200, "account.html", map[string]any{ - "Repo": repo, - "Tokens": tokenInfo, - "flashes": getFlashesFromSession(e, sess), + "Repo": repo, + "Tokens": tokenInfo, + "flashes": getFlashesFromSession(e, sess), + "Accounts": accounts, + "ActiveDid": repo.Repo.Did, }) } diff --git a/server/handle_account_revoke.go b/server/handle_account_revoke.go index c392fc1..c855306 100644 --- a/server/handle_account_revoke.go +++ b/server/handle_account_revoke.go @@ -1,6 +1,8 @@ package server import ( + "errors" + "github.com/haileyok/cocoon/internal/helpers" "github.com/labstack/echo/v4" ) @@ -21,6 +23,9 @@ func (s *Server) handleAccountRevoke(e echo.Context) error { repo, sess, err := s.getSessionRepoOrErr(e) if err != nil { + if !errors.Is(err, ErrSessionUnauthenticated) { + return helpers.ServerError(e, nil) + } return e.Redirect(303, "/account/signin") } diff --git a/server/handle_account_signin.go b/server/handle_account_signin.go index 4490426..91d1b4c 100644 --- a/server/handle_account_signin.go +++ b/server/handle_account_signin.go @@ -1,6 +1,7 @@ package server import ( + "context" "errors" "fmt" "strings" @@ -23,25 +24,51 @@ type OauthSigninInput struct { QueryParams string `form:"query_params"` } -func (s *Server) getSessionRepoOrErr(e echo.Context) (*models.RepoActor, *sessions.Session, error) { - ctx := e.Request().Context() +var ErrSessionUnauthenticated = errors.New("session is unauthenticated") +func (s *Server) getSessionRepoAndAccountsOrErr(e echo.Context) (*models.RepoActor, *sessions.Session, []models.RepoActor, error) { + ctx := e.Request().Context() sess, err := session.Get(s.config.SessionCookieKey, e) if err != nil { - return nil, nil, err + return nil, nil, nil, err } - did, ok := sess.Values["did"].(string) - if !ok { - return nil, sess, errors.New("did was not set in session") + return s.getSessionRepoAndAccountsFromSessionOrErr(e, ctx, sess) +} + +func (s *Server) getSessionRepoAndAccountsFromSessionOrErr(e echo.Context, ctx context.Context, sess *sessions.Session) (*models.RepoActor, *sessions.Session, []models.RepoActor, error) { + if sess == nil { + return nil, nil, nil, errors.New("session is nil") } - repo, err := s.getRepoActorByDid(ctx, did) + accounts, changed, err := s.getSessionAccountActors(ctx, sess) if err != nil { - return nil, sess, err + return nil, sess, nil, err + } + if changed { + applyAccountSessionOptions(sess, int(AccountSessionMaxAge.Seconds())) + if err := sess.Save(e.Request(), e.Response()); err != nil { + return nil, sess, nil, err + } + } + + did := getActiveSessionDid(sess) + if did == "" { + return nil, sess, accounts, fmt.Errorf("%w: did was not set in session", ErrSessionUnauthenticated) + } + + for i := range accounts { + if accounts[i].Repo.Did == did { + return &accounts[i], sess, accounts, nil + } } - return repo, sess, nil + return nil, sess, accounts, fmt.Errorf("%w: did was not found in session accounts", ErrSessionUnauthenticated) +} + +func (s *Server) getSessionRepoOrErr(e echo.Context) (*models.RepoActor, *sessions.Session, error) { + repo, sess, _, err := s.getSessionRepoAndAccountsOrErr(e) + return repo, sess, err } func getFlashesFromSession(e echo.Context, sess *sessions.Session) map[string]any { @@ -54,14 +81,28 @@ func getFlashesFromSession(e echo.Context, sess *sessions.Session) map[string]an } func (s *Server) handleAccountSigninGet(e echo.Context) error { - _, sess, err := s.getSessionRepoOrErr(e) - if err == nil { + repo, sess, accounts, err := s.getSessionRepoAndAccountsOrErr(e) + if err != nil && !errors.Is(err, ErrSessionUnauthenticated) { + return helpers.ServerError(e, nil) + } + if err == nil && e.QueryString() == "" { return e.Redirect(303, "/account") } + if sess == nil { + return helpers.ServerError(e, nil) + } + + activeDid := "" + if repo != nil { + activeDid = repo.Repo.Did + } + return e.Render(200, "signin.html", map[string]any{ "flashes": getFlashesFromSession(e, sess), "QueryParams": e.QueryParams().Encode(), + "Accounts": accounts, + "ActiveDid": activeDid, }) } @@ -161,14 +202,9 @@ func (s *Server) handleAccountSigninPost(e echo.Context) error { } } - sess.Options = &sessions.Options{ - Path: "/", - MaxAge: int(AccountSessionMaxAge.Seconds()), - HttpOnly: true, - } + applyAccountSessionOptions(sess, int(AccountSessionMaxAge.Seconds())) - sess.Values = map[any]any{} - sess.Values["did"] = repo.Repo.Did + setActiveSessionDid(sess, repo.Repo.Did) if err := sess.Save(e.Request(), e.Response()); err != nil { return err diff --git a/server/handle_account_signout.go b/server/handle_account_signout.go index 48e7671..63768f7 100644 --- a/server/handle_account_signout.go +++ b/server/handle_account_signout.go @@ -1,7 +1,6 @@ package server import ( - "github.com/gorilla/sessions" "github.com/labstack/echo-contrib/session" "github.com/labstack/echo/v4" ) @@ -12,13 +11,17 @@ func (s *Server) handleAccountSignout(e echo.Context) error { return err } - sess.Options = &sessions.Options{ - Path: "/", - MaxAge: -1, - HttpOnly: true, + activeDid := getActiveSessionDid(sess) + if activeDid != "" { + removeSessionDid(sess, activeDid) } - sess.Values = map[any]any{} + maxAge := int(AccountSessionMaxAge.Seconds()) + if len(getSessionDids(sess)) == 0 { + maxAge = -1 + } + + applyAccountSessionOptions(sess, maxAge) if err := sess.Save(e.Request(), e.Response()); err != nil { return err diff --git a/server/handle_account_switch.go b/server/handle_account_switch.go new file mode 100644 index 0000000..ac1844d --- /dev/null +++ b/server/handle_account_switch.go @@ -0,0 +1,116 @@ +package server + +import ( + "net/http" + "net/url" + "slices" + "strings" + + "github.com/Azure/go-autorest/autorest/to" + "github.com/haileyok/cocoon/internal/helpers" + "github.com/labstack/echo-contrib/session" + "github.com/labstack/echo/v4" +) + +type AccountSwitchRequest struct { + Did string `form:"did"` + QueryParams string `form:"query_params"` + Next string `form:"next"` +} + +func sanitizeLocalRedirectPath(next string) string { + redirect := strings.TrimSpace(next) + if redirect == "" { + return "/account" + } + if !strings.HasPrefix(redirect, "/") || strings.HasPrefix(redirect, "//") { + return "/account" + } + + parsed, err := url.Parse(redirect) + if err != nil || parsed.IsAbs() || parsed.Host != "" { + return "/account" + } + + return redirect +} + +func mergeRedirectQuery(redirect string, queryParams string) (string, error) { + parsedRedirect, err := url.Parse(redirect) + if err != nil { + return "", err + } + + merged := parsedRedirect.Query() + + rawQueryParams := strings.TrimSpace(queryParams) + if rawQueryParams != "" { + rawQueryParams = strings.TrimPrefix(rawQueryParams, "?") + additional, err := url.ParseQuery(rawQueryParams) + if err != nil { + return "", err + } + for key, values := range additional { + for _, value := range values { + merged.Add(key, value) + } + } + } + + parsedRedirect.RawQuery = merged.Encode() + return parsedRedirect.String(), nil +} + +func isSameOriginRequest(e echo.Context) bool { + host := e.Request().Host + + origin := strings.TrimSpace(e.Request().Header.Get("Origin")) + if origin != "" { + parsedOrigin, err := url.Parse(origin) + return err == nil && parsedOrigin.Host == host + } + + referer := strings.TrimSpace(e.Request().Header.Get("Referer")) + if referer != "" { + parsedReferer, err := url.Parse(referer) + return err == nil && parsedReferer.Host == host + } + + return false +} + +func (s *Server) handleAccountSwitchPost(e echo.Context) error { + if !isSameOriginRequest(e) { + return e.JSON(http.StatusForbidden, map[string]string{"error": "Forbidden"}) + } + + var req AccountSwitchRequest + if err := e.Bind(&req); err != nil { + return helpers.InputError(e, to.StringPtr("invalid switch account request")) + } + + sess, err := session.Get(s.config.SessionCookieKey, e) + if err != nil { + return err + } + + dids := getSessionDids(sess) + if !slices.Contains(dids, req.Did) { + return helpers.InputError(e, to.StringPtr("requested account is not logged in")) + } + + setActiveSessionDid(sess, req.Did) + applyAccountSessionOptions(sess, int(AccountSessionMaxAge.Seconds())) + + if err := sess.Save(e.Request(), e.Response()); err != nil { + return err + } + + redirect := sanitizeLocalRedirectPath(req.Next) + redirect, err = mergeRedirectQuery(redirect, req.QueryParams) + if err != nil { + return helpers.InputError(e, to.StringPtr("invalid query params")) + } + + return e.Redirect(303, redirect) +} diff --git a/server/handle_oauth_authorize.go b/server/handle_oauth_authorize.go index 2665c7a..e1834e9 100644 --- a/server/handle_oauth_authorize.go +++ b/server/handle_oauth_authorize.go @@ -1,8 +1,10 @@ package server import ( + "errors" "fmt" "net/url" + "slices" "strings" "time" @@ -11,6 +13,7 @@ import ( "github.com/haileyok/cocoon/oauth" "github.com/haileyok/cocoon/oauth/constants" "github.com/haileyok/cocoon/oauth/provider" + "github.com/labstack/echo-contrib/session" "github.com/labstack/echo/v4" ) @@ -52,6 +55,8 @@ func (s *Server) handleOauthAuthorizeGet(e echo.Context) error { "AppName": "DEV MODE AUTHORIZATION PAGE", "Handle": "paula.cocoon.social", "RequestUri": "", + "Accounts": []string{}, + "ActiveDid": "", }) } return helpers.InputError(e, to.StringPtr("no request uri and invalid parameters")) @@ -96,11 +101,6 @@ func (s *Server) handleOauthAuthorizeGet(e echo.Context) error { } - repo, _, err := s.getSessionRepoOrErr(e) - if err != nil { - return e.Redirect(303, "/account/signin?"+e.QueryParams().Encode()) - } - var req provider.OauthAuthorizationRequest if err := s.db.Raw(ctx, "SELECT * FROM oauth_authorization_requests WHERE request_id = ?", nil, reqId).Scan(&req).Error; err != nil { return helpers.ServerError(e, to.StringPtr(err.Error())) @@ -116,6 +116,32 @@ func (s *Server) handleOauthAuthorizeGet(e echo.Context) error { return helpers.ServerError(e, to.StringPtr(err.Error())) } + sess, err := session.Get(s.config.SessionCookieKey, e) + if err != nil { + return helpers.ServerError(e, to.StringPtr(err.Error())) + } + + if req.Parameters.LoginHint != nil && *req.Parameters.LoginHint != "" { + did, err := s.resolveLoginHintToDid(ctx, *req.Parameters.LoginHint) + if err != nil || !slices.Contains(getSessionDids(sess), did) { + return e.Redirect(303, "/account/signin?"+e.QueryParams().Encode()) + } + + setActiveSessionDid(sess, did) + applyAccountSessionOptions(sess, int(AccountSessionMaxAge.Seconds())) + if err := sess.Save(e.Request(), e.Response()); err != nil { + return helpers.ServerError(e, to.StringPtr(err.Error())) + } + } + + repo, _, accounts, err := s.getSessionRepoAndAccountsFromSessionOrErr(e, ctx, sess) + if err != nil { + if !errors.Is(err, ErrSessionUnauthenticated) { + return helpers.ServerError(e, to.StringPtr(err.Error())) + } + return e.Redirect(303, "/account/signin?"+e.QueryParams().Encode()) + } + scopes := strings.Split(req.Parameters.Scope, " ") appName := client.Metadata.ClientName @@ -125,6 +151,8 @@ func (s *Server) handleOauthAuthorizeGet(e echo.Context) error { "RequestUri": input.RequestUri, "QueryParams": e.QueryParams().Encode(), "Handle": repo.Actor.Handle, + "Accounts": accounts, + "ActiveDid": repo.Repo.Did, } return e.Render(200, "authorize.html", data) @@ -141,6 +169,9 @@ func (s *Server) handleOauthAuthorizePost(e echo.Context) error { repo, _, err := s.getSessionRepoOrErr(e) if err != nil { + if !errors.Is(err, ErrSessionUnauthenticated) { + return helpers.ServerError(e, to.StringPtr(err.Error())) + } return e.Redirect(303, "/account/signin") } diff --git a/server/server.go b/server/server.go index f92b15b..2a519d1 100644 --- a/server/server.go +++ b/server/server.go @@ -523,6 +523,7 @@ func (s *Server) addRoutes() { // account s.echo.GET("/account", s.handleAccount) s.echo.POST("/account/revoke", s.handleAccountRevoke) + s.echo.POST("/account/switch", s.handleAccountSwitchPost) s.echo.GET("/account/signin", s.handleAccountSigninGet) s.echo.POST("/account/signin", s.handleAccountSigninPost) s.echo.GET("/account/signout", s.handleAccountSignout) diff --git a/server/session_options.go b/server/session_options.go new file mode 100644 index 0000000..f1d4581 --- /dev/null +++ b/server/session_options.go @@ -0,0 +1,16 @@ +package server + +import ( + "net/http" + + "github.com/gorilla/sessions" +) + +func applyAccountSessionOptions(sess *sessions.Session, maxAge int) { + sess.Options = &sessions.Options{ + Path: "/", + MaxAge: maxAge, + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + } +} diff --git a/server/templates/account.html b/server/templates/account.html index 6035679..162f79e 100644 --- a/server/templates/account.html +++ b/server/templates/account.html @@ -12,8 +12,23 @@

Welcome, {{ .Repo.Handle }}

+ {{ if gt (len .Accounts) 1 }} +
+ + + + +
+ {{ end }} {{ if .flashes.successes }}

{{ index .flashes.successes 0 }}

diff --git a/server/templates/authorize.html b/server/templates/authorize.html index 099b9c7..6c3180d 100644 --- a/server/templates/authorize.html +++ b/server/templates/authorize.html @@ -15,7 +15,25 @@

Authorizing with {{ .AppName }}

You are signed in as {{ .Handle }}. - Switch Account + Sign out this account +

+ {{ if gt (len .Accounts) 1 }} +
+ + + + + +
+ {{ end }} +

+ Need a different account? Sign in another account.

{{ .AppName }} is asking for you to grant it these scopes:

    diff --git a/server/templates/signin.html b/server/templates/signin.html index 3348c26..c025e3d 100644 --- a/server/templates/signin.html +++ b/server/templates/signin.html @@ -12,6 +12,9 @@

    Sign into your account

    Enter your handle and password below.

    + {{ if gt (len .Accounts) 0 }} +

    You currently have {{ len .Accounts }} signed-in account(s).

    + {{ end }} {{ if .flashes.errors }}

    {{ index .flashes.errors 0 }}

    diff --git a/test.go b/test.go index 5e9fa46..e2d3041 100644 --- a/test.go +++ b/test.go @@ -10,10 +10,10 @@ import ( "strings" "github.com/bluesky-social/indigo/api/atproto" + atp "github.com/bluesky-social/indigo/atproto/repo" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/bluesky-social/indigo/events" "github.com/bluesky-social/indigo/events/schedulers/parallel" - atp "github.com/bluesky-social/indigo/atproto/repo" lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/bluesky-social/indigo/repomgr" "github.com/gorilla/websocket" -- 2.51.2