middleware_test.go
1package rtw
2
3import (
4 "net/http"
5 "net/http/httptest"
6 "os"
7 "testing"
8 "time"
9
10 "git.kilimanjaro.io/rtw/httperr"
11 "git.kilimanjaro.io/rtw/pkg/id"
12 "git.kilimanjaro.io/rtw/pkg/session"
13 "git.kilimanjaro.io/rtw/user"
14 "github.com/stretchr/testify/assert"
15 "github.com/stretchr/testify/require"
16)
17
18func setupTestMiddleware(t *testing.T) (*user.Service, *httperr.Handler, func()) {
19 tmpDir, err := os.MkdirTemp("", "middleware_test_*")
20 require.NoError(t, err)
21
22 us, err := user.NewService(user.WithBasePath(tmpDir))
23 require.NoError(t, err)
24
25 eh := httperr.NewHandler()
26
27 cleanup := func() {
28 us.Shutdown()
29 os.RemoveAll(tmpDir)
30 }
31
32 return us, eh, cleanup
33}
34
35func mockNextHandler(captureUserID *id.Key, captureUserType *session.UserType) http.Handler {
36 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
37 info, err := session.FromContext[id.Key](r.Context())
38 if err == nil {
39 *captureUserID = info.ID
40 *captureUserType = info.Type
41 }
42 w.WriteHeader(http.StatusOK)
43 w.Write([]byte("OK"))
44 })
45}
46
47func TestSession_NoCookie_CreatesAnonymousSession(t *testing.T) {
48 us, eh, cleanup := setupTestMiddleware(t)
49 defer cleanup()
50
51 var capturedUserID id.Key
52 var capturedUserType session.UserType
53
54 middleware := Session(us, eh, CreateAnonymousSessions(true))
55 handler := middleware(mockNextHandler(&capturedUserID, &capturedUserType))
56
57 req := httptest.NewRequest(http.MethodGet, "/test", nil)
58 rr := httptest.NewRecorder()
59
60 handler.ServeHTTP(rr, req)
61
62 // Assert: Should return 200 (next handler called)
63 assert.Equal(t, http.StatusOK, rr.Code)
64
65 // Assert: User ID should be set (anonymous session created)
66 assert.NotEqual(t, id.NotExist, capturedUserID)
67 assert.NotEqual(t, id.Key(0), capturedUserID)
68
69 // Assert: User type should be Anonymous
70 assert.Equal(t, session.Anonymous, capturedUserType)
71
72 // Assert: Cookie should be set in response
73 cookies := rr.Result().Cookies()
74 require.Len(t, cookies, 1)
75 assert.Equal(t, "__Host-session-id", cookies[0].Name)
76 assert.NotEmpty(t, cookies[0].Value)
77 assert.True(t, cookies[0].HttpOnly)
78 assert.True(t, cookies[0].Secure)
79 assert.Equal(t, http.SameSiteStrictMode, cookies[0].SameSite)
80}
81
82func TestSession_ExpiredSession_RedirectsToLogin(t *testing.T) {
83 us, eh, cleanup := setupTestMiddleware(t)
84 defer cleanup()
85
86 // Create a user and session first
87 email := "[email protected]"
88 password := "password123"
89 userID, err := us.Signup(&email, nil, password)
90 require.NoError(t, err)
91
92 // Create a session that is already expired
93 sessionID := session.NewID()
94 err = session.Persist(us.DB(), sessionID, userID, time.Now().UTC().Add(-1*time.Hour))
95 require.NoError(t, err)
96
97 var capturedUserID id.Key
98 var capturedUserType session.UserType
99
100 middleware := Session(us, eh)
101 handler := middleware(mockNextHandler(&capturedUserID, &capturedUserType))
102
103 req := httptest.NewRequest(http.MethodGet, "/test", nil)
104 req.AddCookie(&http.Cookie{
105 Name: "__Host-session-id",
106 Value: sessionID,
107 })
108 rr := httptest.NewRecorder()
109
110 handler.ServeHTTP(rr, req)
111
112 // Assert: Should redirect to login page
113 assert.Equal(t, http.StatusFound, rr.Code)
114
115 // Assert: Location header should be /login
116 location := rr.Header().Get("Location")
117 assert.Equal(t, "/login?next=%2Ftest", location)
118
119 // Assert: Next handler should NOT be called (user ID should be zero)
120 assert.Equal(t, id.Key(0), capturedUserID)
121}
122
123func TestSession_ValidSession_SetsUserIDInContext(t *testing.T) {
124 us, eh, cleanup := setupTestMiddleware(t)
125 defer cleanup()
126
127 // Create a registered user
128 email := "[email protected]"
129 password := "password123"
130 userID, err := us.Signup(&email, nil, password)
131 require.NoError(t, err)
132
133 // Create a valid session
134 sessionID := session.NewID()
135 err = session.Persist(us.DB(), sessionID, userID, time.Now().UTC().Add(24*time.Hour))
136 require.NoError(t, err)
137
138 var capturedUserID id.Key
139 var capturedUserType session.UserType
140
141 middleware := Session(us, eh)
142 handler := middleware(mockNextHandler(&capturedUserID, &capturedUserType))
143
144 req := httptest.NewRequest(http.MethodGet, "/test", nil)
145 req.AddCookie(&http.Cookie{
146 Name: "__Host-session-id",
147 Value: sessionID,
148 })
149 rr := httptest.NewRecorder()
150
151 handler.ServeHTTP(rr, req)
152
153 // Assert: Should return 200 (next handler called)
154 assert.Equal(t, http.StatusOK, rr.Code)
155
156 // Assert: User ID should match the registered user
157 assert.Equal(t, userID, capturedUserID)
158
159 // Assert: User type should be Registered
160 assert.Equal(t, session.Registered, capturedUserType)
161}