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}