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}