middleware.go

  1package rtw
  2
  3import (
  4	"errors"
  5	"fmt"
  6	"net/http"
  7
  8	"git.kilimanjaro.io/rtw/httperr"
  9	"git.kilimanjaro.io/rtw/pkg/id"
 10	"git.kilimanjaro.io/rtw/pkg/log"
 11	"git.kilimanjaro.io/rtw/pkg/middleware"
 12	"git.kilimanjaro.io/rtw/pkg/session"
 13	"git.kilimanjaro.io/rtw/user"
 14)
 15
 16type sessionConfig struct {
 17	useAnonymous      bool
 18	noExpiredRedirect bool
 19}
 20
 21type SessionOption func(*sessionConfig)
 22
 23func CreateAnonymousSessions(should bool) SessionOption {
 24	return func(s *sessionConfig) {
 25		s.useAnonymous = should
 26	}
 27}
 28
 29// NoExpiredRedirect prevents redirecting to the login page when the session
 30// cookie exists but the session is expired or not found in the database.
 31// Instead, the expired cookie is cleared and the request passes through to the
 32// next handler. This is intended for the /login and /signup routes to prevent
 33// redirect loops when a user visits them with a stale session cookie.
 34func NoExpiredRedirect() SessionOption {
 35	return func(s *sessionConfig) {
 36		s.noExpiredRedirect = true
 37	}
 38}
 39
 40func Session(u *user.Service, eh *httperr.Handler, opts ...SessionOption) middleware.Middleware {
 41	cfg := &sessionConfig{
 42		useAnonymous: false,
 43	}
 44	for _, opt := range opts {
 45		opt(cfg)
 46	}
 47	return func(next http.Handler) http.Handler {
 48		return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
 49			l := log.FromContext(r.Context())
 50			sessionID, err := session.GetCookie(r)
 51			if err != nil {
 52				if cfg.useAnonymous {
 53					l.Debug("no session, creating new anon session")
 54					// create anon session if no cookie present
 55					anonID, err := u.CreateNewAnonSession(w)
 56					if err != nil {
 57						l.Error("failed to create new anon session", "error", err)
 58						eh.Handle(http.StatusInternalServerError, w, r, fmt.Errorf("system error: %w", err))
 59						return
 60					}
 61					ctx := session.ToContext(r.Context(), session.NewInfo(anonID, session.Anonymous))
 62					next.ServeHTTP(w, r.WithContext(ctx))
 63				} else {
 64					// no automatic anon session, pass through to next handler
 65					next.ServeHTTP(w, r)
 66				}
 67				return
 68			}
 69			userID, err := session.Read[id.Key](u.DB(), sessionID)
 70			if err != nil {
 71				if errors.Is(err, session.ErrNoSession) {
 72					if cfg.noExpiredRedirect {
 73						l.Debug("session expired, clearing cookie and passing through")
 74						session.ExpireCookie(w)
 75						next.ServeHTTP(w, r)
 76						return
 77					}
 78					l.Error("session expired", "user", userID)
 79					eh.Redirect("/login", w, r, "next", r.URL.Path)
 80					return
 81				}
 82				l.Error("failed to read session from database", "error", err)
 83				eh.Handle(http.StatusInternalServerError, w, r, fmt.Errorf("system error: %w", err))
 84				return
 85			}
 86			user, err := u.GetUserByID(userID)
 87			if err != nil {
 88				l.Error("failed to get user by ID", "error", err)
 89				eh.Handle(http.StatusInternalServerError, w, r, fmt.Errorf("system error: %w", err))
 90				return
 91			}
 92			var userType session.UserType
 93			if user.Registered {
 94				userType = session.Registered
 95			} else {
 96				userType = session.Anonymous
 97			}
 98			l.Debug("session exists", "user_id", userID, "user_type", userType)
 99			ctx := session.ToContext(r.Context(), session.NewInfo(userID, userType))
100			next.ServeHTTP(w, r.WithContext(ctx))
101		})
102	}
103}