diff --git a/README.md b/README.md index 996e9b1..d1c46bd 100644 --- a/README.md +++ b/README.md @@ -1,31 +1,39 @@ # go-session-redis Simple HTTP session management -## Example middleware (Chi router) +## Example middleware ```go -func SessionAuthenticator(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - s, err := session.SessionManager.GetOrCreateSession(w, r) - if err != nil { - http.Error(w, err.Error(), http.StatusUnauthorized) - return - } +func SessionAuthenticator(manager *session.Manager, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + s, err := manager.GetSession(w, r) + if errors.Is(err, session.ErrSessionNotFound) { + http.Error(w, "Not authenticated", http.StatusUnauthorized) + return + } + if err != nil { + http.Error(w, err.Error(), http.StatusBadGateway) + return + } - id := s.Get("id") - if id == "" || id == nil { - http.Error(w, "Not authenticated", http.StatusUnauthorized) - return - } + userID, ok := s.Get("user_id").(string) + if !ok || userID == "" { + http.Error(w, "Not authenticated", http.StatusUnauthorized) + return + } - err = s.UpdateTTL() - if err != nil { - http.Error(w, err.Error(), http.StatusBadGateway) - return - } + // GetSession refreshes the TTL. + next.ServeHTTP(w, r) + }) +} +``` - // Session is valid - next.ServeHTTP(w, r) - }) +## Creating a session +```go +func CreateAuthenticatedSession(manager *session.Manager, w http.ResponseWriter, r *http.Request, userID string) error { + s, err := manager.CreateSession(w, r) + if err != nil { + return err + } + return s.Set("user_id", userID) } - ``` diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..45ecb83 --- /dev/null +++ b/go.mod @@ -0,0 +1,15 @@ +module github.com/Angel-Technologies/go-session-redis + +go 1.27 + +require ( + github.com/alicebob/miniredis/v2 v2.39.0 + github.com/redis/go-redis/v9 v9.22.0 +) + +require ( + github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/yuin/gopher-lua v1.1.1 // indirect + go.uber.org/atomic v1.11.0 // indirect + golang.org/x/sys v0.30.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..c41bee7 --- /dev/null +++ b/go.sum @@ -0,0 +1,26 @@ +github.com/alicebob/miniredis/v2 v2.39.0 h1:M7WbmV5BmV56L8KTG0rw6vEQ+woTOghpDgin2xv4A0g= +github.com/alicebob/miniredis/v2 v2.39.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= +github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= +github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c= +github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= +github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE= +github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/redis/go-redis/v9 v9.22.0 h1:laDvpYXTJtZLloinw1fA5Kqd6HAEH2XKxOkG/PDq2F0= +github.com/redis/go-redis/v9 v9.22.0/go.mod h1:y2g0Wj8rQvuK0ELM+oxSudcLtC09JScs98I/X9gRWY4= +github.com/stretchr/testify v1.3.0 h1:TivCn/peBQ7UY8ooIcPgZFpTNSz0Q2U6UrFlUfqbe0Q= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= +github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= +github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= +github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= +golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc= +golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= diff --git a/session_test.go b/session_test.go new file mode 100644 index 0000000..8ead19f --- /dev/null +++ b/session_test.go @@ -0,0 +1,1137 @@ +package session + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "sync" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" +) + +// redisTestHook can pause one Redis command and/or inject an error into script +// and TTL commands. The optional entered channel lets a test wait until the +// command has reached the pause point without relying on a sleep. +type redisTestHook struct { + gateCommand string + commandGate chan struct{} + gateUsed bool + commandErrs map[string]error + + mu sync.Mutex + entered chan struct{} + enteredOnce sync.Once +} + +func (hook *redisTestHook) DialHook(next redis.DialHook) redis.DialHook { + return next +} + +func (hook *redisTestHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, command redis.Cmder) error { + commandName := strings.ToLower(command.Name()) + + hook.mu.Lock() + gate := (<-chan struct{})(nil) + if commandName == hook.gateCommand && hook.commandGate != nil && !hook.gateUsed { + hook.gateUsed = true + gate = hook.commandGate + } + injectedErr := hook.commandErrs[commandName] + hook.mu.Unlock() + + if gate != nil { + if hook.entered != nil { + hook.enteredOnce.Do(func() { close(hook.entered) }) + } + <-gate + } + if injectedErr != nil && (commandName == "evalsha" || commandName == "eval" || commandName == "pttl") { + return injectedErr + } + return next(ctx, command) + } +} + +func (hook *redisTestHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return next +} + +type redisContextCaptureHook struct { + mu sync.Mutex + contexts []context.Context +} + +func (hook *redisContextCaptureHook) DialHook(next redis.DialHook) redis.DialHook { + return next +} + +func (hook *redisContextCaptureHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, command redis.Cmder) error { + commandName := strings.ToLower(command.Name()) + if commandName == "evalsha" || commandName == "eval" { + hook.mu.Lock() + hook.contexts = append(hook.contexts, ctx) + hook.mu.Unlock() + } + return next(ctx, command) + } +} + +func (hook *redisContextCaptureHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return next +} + +func (hook *redisContextCaptureHook) snapshot() []context.Context { + hook.mu.Lock() + defer hook.mu.Unlock() + return append([]context.Context(nil), hook.contexts...) +} + +func newTest(t *testing.T, ttl time.Duration) (*miniredis.Miniredis, *redis.Client, *Manager) { + t.Helper() + miniRedis := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: miniRedis.Addr()}) + t.Cleanup(func() { _ = client.Close() }) + return miniRedis, client, NewManager(client, "sid", ttl, false) +} + +func (manager *Manager) cachedSession(sid string) (*Session, bool) { + manager.cacheMu.Lock() + defer manager.cacheMu.Unlock() + session, found := manager.cache[sid] + return session, found +} + +func seedCachedSession(manager *Manager, session *Session) *Session { + manager.cacheMu.Lock() + defer manager.cacheMu.Unlock() + manager.cache[session.sid] = session + return session +} + +func expireSessionForTest(session *Session) { + session.mutex.Lock() + session.expiresAt = time.Now().Add(-time.Second) + session.mutex.Unlock() + session.expire() +} + +func getSessionRequest(t *testing.T, manager *Manager, sid string) (SessionInterface, *httptest.ResponseRecorder, error) { + t.Helper() + if sid == "" { + t.Fatal("getSessionRequest requires a non-empty session ID") + } + request := httptest.NewRequest(http.MethodGet, "/", nil) + request.AddCookie(&http.Cookie{Name: "sid", Value: sid}) + recorder := httptest.NewRecorder() + session, err := manager.GetSession(recorder, request) + return session, recorder, err +} + +func responseSessionID(t *testing.T, recorder *httptest.ResponseRecorder) string { + t.Helper() + cookies := recorder.Result().Cookies() + if len(cookies) == 0 { + return "" + } + return cookies[0].Value +} + +func createSession(t *testing.T, manager *Manager) (*Session, string) { + t.Helper() + request := httptest.NewRequest(http.MethodGet, "/", nil) + recorder := httptest.NewRecorder() + sessionInterface, err := manager.CreateSession(recorder, request) + if err != nil { + t.Fatalf("CreateSession returned an error: %v", err) + } + session, ok := sessionInterface.(*Session) + if !ok { + t.Fatalf("CreateSession returned %T, want *Session", sessionInterface) + } + return session, responseSessionID(t, recorder) +} + +func waitForHookEntry(t *testing.T, entered <-chan struct{}) { + t.Helper() + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("timed out waiting for Redis command to reach test hook") + } +} + +func TestCreateSessionCanceledRequest(t *testing.T) { + _, _, manager := newTest(t, time.Minute) + requestContext, cancel := context.WithCancel(context.Background()) + cancel() + request := httptest.NewRequest(http.MethodGet, "/", nil).WithContext(requestContext) + recorder := httptest.NewRecorder() + + session, err := manager.CreateSession(recorder, request) + if session != nil { + t.Fatalf("canceled request returned session %v, want nil", session) + } + if !errors.Is(err, context.Canceled) { + t.Fatalf("canceled request error = %v, want context.Canceled", err) + } + if cookies := recorder.Result().Cookies(); len(cookies) != 0 { + t.Fatalf("canceled request emitted %d cookies, want none", len(cookies)) + } +} + +func TestGetSessionRequestDeadline(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + _, sessionID := createSession(t, manager) + manager = NewManager(client, "sid", time.Minute, false) + contextHook := &redisContextCaptureHook{} + client.AddHook(contextHook) + + request := httptest.NewRequest(http.MethodGet, "/", nil) + request.AddCookie(&http.Cookie{Name: "sid", Value: sessionID}) + recorder := httptest.NewRecorder() + startedAt := time.Now() + session, err := manager.GetSession(recorder, request) + if err != nil { + t.Fatalf("GetSession returned an error: %v", err) + } + if session == nil { + t.Fatal("GetSession returned nil session, want existing session") + } + + contexts := contextHook.snapshot() + if len(contexts) < 2 { + t.Fatalf("captured %d Redis script contexts, want read and refresh contexts", len(contexts)) + } + firstDeadline, hasDeadline := contexts[0].Deadline() + if !hasDeadline { + t.Fatal("first Redis script context has no deadline") + } + if firstDeadline.Before(startedAt.Add(800*time.Millisecond)) || firstDeadline.After(startedAt.Add(1100*time.Millisecond)) { + t.Fatalf("first Redis deadline = %v, want approximately one second after request start %v", firstDeadline, startedAt) + } + for index, redisContext := range contexts[1:] { + deadline, hasDeadline := redisContext.Deadline() + if !hasDeadline || !deadline.Equal(firstDeadline) { + t.Fatalf("Redis script context %d deadline = %v (has deadline = %v), want shared deadline %v", index+1, deadline, hasDeadline, firstDeadline) + } + if redisContext != contexts[0] { + t.Fatalf("Redis script context %d is not the request-derived context", index+1) + } + } +} + +func TestCreateSessionIgnoresExistingCookie(t *testing.T) { + miniRedis, _, manager := newTest(t, time.Minute) + oldSession, oldSessionID := createSession(t, manager) + if err := oldSession.Set("marker", "old"); err != nil { + t.Fatalf("setting marker on old session: %v", err) + } + + request := httptest.NewRequest(http.MethodGet, "/", nil) + request.AddCookie(&http.Cookie{Name: "sid", Value: oldSessionID}) + recorder := httptest.NewRecorder() + createdSession, err := manager.CreateSession(recorder, request) + if err != nil { + t.Fatalf("CreateSession returned an error: %v", err) + } + newSessionID := responseSessionID(t, recorder) + if createdSession == nil { + t.Fatal("CreateSession returned nil session") + } + if newSessionID == "" { + t.Fatal("CreateSession emitted an empty session ID") + } + if createdSession.ID() != newSessionID { + t.Fatalf("created session ID = %q, response cookie ID = %q", createdSession.ID(), newSessionID) + } + if newSessionID == oldSessionID { + t.Fatalf("CreateSession reused old session ID %q", oldSessionID) + } + oldKey := sessionKeyPrefix + oldSessionID + if !miniRedis.Exists(oldKey) || miniRedis.HGet(oldKey, encodeUserField("marker")) != `"old"` { + t.Fatal("creating a new session changed or removed the old Redis session") + } +} + +func TestGetSessionWithoutID(t *testing.T) { + testCases := []struct { + name string + cookie *http.Cookie + }{ + {name: "no cookie"}, + {name: "empty cookie", cookie: &http.Cookie{Name: "sid", Value: ""}}, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + request := httptest.NewRequest(http.MethodGet, "/", nil) + if testCase.cookie != nil { + request.AddCookie(testCase.cookie) + } + recorder := httptest.NewRecorder() + + gotSession, gotErr := manager.GetSession(recorder, request) + if gotSession != nil { + t.Fatalf("GetSession session = %v, want nil", gotSession) + } + if !errors.Is(gotErr, ErrSessionNotFound) { + t.Fatalf("GetSession error = %v, want ErrSessionNotFound", gotErr) + } + if gotCookies := len(recorder.Result().Cookies()); gotCookies != 0 { + t.Fatalf("missing session lookup cookies = %d, want 0", gotCookies) + } + if gotCacheSize := len(manager.cache); gotCacheSize != 0 { + t.Fatalf("missing session lookup cache size = %d, want 0", gotCacheSize) + } + keys, err := client.Keys(context.Background(), sessionKeyPrefix+"*").Result() + if err != nil { + t.Fatalf("listing session keys: %v", err) + } + if len(keys) != 0 { + t.Fatalf("missing session lookup created Redis keys: %v", keys) + } + }) + } +} + +func TestGetSessionRefreshesTTL(t *testing.T) { + const managerTTL = time.Minute + miniRedis, client, manager := newTest(t, managerTTL) + createdSession, sessionID := createSession(t, manager) + miniRedis.FastForward(20 * time.Second) + beforeRefresh, err := client.PTTL(context.Background(), sessionKeyPrefix+sessionID).Result() + if err != nil { + t.Fatalf("reading TTL before retrieval: %v", err) + } + if beforeRefresh >= managerTTL-time.Second { + t.Fatalf("TTL before retrieval = %v, want less than configured TTL %v", beforeRefresh, managerTTL) + } + + foundSession, recorder, err := getSessionRequest(t, manager, sessionID) + if err != nil { + t.Fatalf("GetSession returned an error: %v", err) + } + wantSession := createdSession + if gotSession := foundSession; gotSession != wantSession { + t.Fatalf("GetSession returned %p, want interned session %p", gotSession, wantSession) + } + wantSessionID := sessionID + if gotSessionID := foundSession.ID(); gotSessionID != wantSessionID { + t.Fatalf("retrieved session ID = %q, want %q", gotSessionID, wantSessionID) + } + wantCookieID := sessionID + if gotCookieID := responseSessionID(t, recorder); gotCookieID != wantCookieID { + t.Fatalf("response session ID = %q, want %q", gotCookieID, wantCookieID) + } + afterRefresh, err := client.PTTL(context.Background(), sessionKeyPrefix+sessionID).Result() + if err != nil { + t.Fatalf("reading TTL after retrieval: %v", err) + } + wantMinTTL, wantMaxTTL := managerTTL-time.Second, managerTTL + if afterRefresh < wantMinTTL || afterRefresh > wantMaxTTL { + t.Fatalf("TTL after retrieval = %v, want within one second of configured TTL %v", afterRefresh, managerTTL) + } +} + +func TestGetSessionConcurrentLookupsShareSession(t *testing.T) { + miniRedis, client, writerManager := newTest(t, time.Minute) + _, sessionID := createSession(t, writerManager) + readerManager := NewManager(client, "sid", time.Minute, false) + readGate := make(chan struct{}) + readEntered := make(chan struct{}) + client.AddHook(&redisTestHook{ + gateCommand: "evalsha", + commandGate: readGate, + entered: readEntered, + }) + + const callerCount = 2 + start := make(chan struct{}) + results := make(chan struct { + session SessionInterface + err error + }, callerCount) + var waitGroup sync.WaitGroup + waitGroup.Add(callerCount) + for range callerCount { + go func() { + defer waitGroup.Done() + <-start + session, _, err := getSessionRequest(t, readerManager, sessionID) + results <- struct { + session SessionInterface + err error + }{session: session, err: err} + }() + } + close(start) + waitForHookEntry(t, readEntered) + close(readGate) + waitGroup.Wait() + close(results) + + var sessions []SessionInterface + for result := range results { + if result.err != nil { + t.Fatalf("concurrent GetSession returned an error: %v", result.err) + } + sessions = append(sessions, result.session) + } + if len(sessions) != callerCount { + t.Fatalf("received %d concurrent results, want %d", len(sessions), callerCount) + } + if sessions[0].(*Session) != sessions[1].(*Session) { + t.Fatal("concurrent lookups returned different in-memory sessions") + } + if len(readerManager.cache) != 1 { + t.Fatalf("reader cache contains %d sessions, want one interned entry", len(readerManager.cache)) + } + if !miniRedis.Exists(sessionKeyPrefix + sessionID) { + t.Fatal("concurrent lookup removed the Redis session") + } +} + +func TestGetSessionMissingSIDDoesNotCreate(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + missingSessionID := "expired-concurrent" + seedCachedSession(manager, &Session{sid: missingSessionID, manager: manager}) + readGate := make(chan struct{}) + readEntered := make(chan struct{}) + client.AddHook(&redisTestHook{ + gateCommand: "evalsha", + commandGate: readGate, + entered: readEntered, + }) + + const callerCount = 8 + type lookupResult struct { + session SessionInterface + cookies int + err error + } + start := make(chan struct{}) + ready := make(chan struct{}, callerCount) + results := make(chan lookupResult, callerCount) + var waitGroup sync.WaitGroup + waitGroup.Add(callerCount) + for range callerCount { + go func() { + defer waitGroup.Done() + <-start + ready <- struct{}{} + session, recorder, err := getSessionRequest(t, manager, missingSessionID) + results <- lookupResult{ + session: session, + cookies: len(recorder.Result().Cookies()), + err: err, + } + }() + } + close(start) + for range callerCount { + <-ready + } + waitForHookEntry(t, readEntered) + close(readGate) + waitGroup.Wait() + close(results) + + for result := range results { + if result.session != nil { + t.Fatalf("missing lookup returned session %v, want nil", result.session) + } + if !errors.Is(result.err, ErrSessionNotFound) { + t.Fatalf("missing lookup error = %v, want ErrSessionNotFound", result.err) + } + if result.cookies != 0 { + t.Fatalf("missing lookup emitted %d cookies, want none", result.cookies) + } + } + keys, err := client.Keys(context.Background(), sessionKeyPrefix+"*").Result() + if err != nil { + t.Fatalf("listing session keys: %v", err) + } + if len(keys) != 0 { + t.Fatalf("missing lookup created Redis keys: %v", keys) + } + if _, found := manager.cachedSession(missingSessionID); found { + t.Fatal("missing lookup retained a cache placeholder") + } +} + +func TestRedisReadSerializesMutation(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + readGate := make(chan struct{}) + readEntered := make(chan struct{}) + client.AddHook(&redisTestHook{ + gateCommand: "evalsha", + commandGate: readGate, + entered: readEntered, + }) + + readDone := make(chan struct{}) + go func() { + _, _, _ = getSessionRequest(t, manager, sessionID) + close(readDone) + }() + waitForHookEntry(t, readEntered) + + mutationDone := make(chan error) + go func() { mutationDone <- session.Set("x", 1) }() + select { + case err := <-mutationDone: + t.Fatalf("session mutation completed before Redis read released: %v", err) + case <-time.After(20 * time.Millisecond): + } + close(readGate) + if err := <-mutationDone; err != nil { + t.Fatalf("session mutation after read release: %v", err) + } + <-readDone +} + +func TestDifferentSessionsIndependentLocks(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + _, firstSessionID := createSession(t, manager) + _, secondSessionID := createSession(t, manager) + readGate := make(chan struct{}) + readEntered := make(chan struct{}) + client.AddHook(&redisTestHook{ + gateCommand: "evalsha", + commandGate: readGate, + entered: readEntered, + }) + + firstReadDone := make(chan error) + go func() { + _, _, err := getSessionRequest(t, manager, firstSessionID) + firstReadDone <- err + }() + waitForHookEntry(t, readEntered) + if _, _, err := getSessionRequest(t, manager, secondSessionID); err != nil { + t.Fatalf("lookup of independent session blocked or failed: %v", err) + } + close(readGate) + if err := <-firstReadDone; err != nil { + t.Fatalf("first session lookup after release: %v", err) + } +} + +func TestManagersUpdateFieldsWithoutLostWrites(t *testing.T) { + _, client, firstManager := newTest(t, time.Minute) + secondManager := NewManager(client, "sid", time.Minute, false) + _, sessionID := createSession(t, firstManager) + firstSession, _, err := getSessionRequest(t, firstManager, sessionID) + if err != nil { + t.Fatalf("loading first manager session: %v", err) + } + secondSession, _, err := getSessionRequest(t, secondManager, sessionID) + if err != nil { + t.Fatalf("loading second manager session: %v", err) + } + if err := firstSession.Set("a", "one"); err != nil { + t.Fatalf("setting field a: %v", err) + } + if err := secondSession.Set("b", "two"); err != nil { + t.Fatalf("setting field b: %v", err) + } + + values, err := client.HGetAll(context.Background(), sessionKeyPrefix+sessionID).Result() + if err != nil { + t.Fatalf("reading session hash: %v", err) + } + if values[encodeUserField("a")] != `"one"` || values[encodeUserField("b")] != `"two"` { + t.Fatalf("session hash fields = %v, want both independent writes", values) + } + freshSession, _, err := getSessionRequest(t, NewManager(client, "sid", time.Minute, false), sessionID) + if err != nil { + t.Fatalf("loading fresh manager session: %v", err) + } + if got := freshSession.Get("a"); got != "one" { + t.Fatalf("fresh session field a = %v, want %q", got, "one") + } + if got := freshSession.Get("b"); got != "two" { + t.Fatalf("fresh session field b = %v, want %q", got, "two") + } +} + +func TestSessionJSONValuesAreIsolated(t *testing.T) { + _, _, manager := newTest(t, time.Minute) + session, _ := createSession(t, manager) + input := map[string]any{"x": []any{1.0}} + if err := session.Set("obj", input); err != nil { + t.Fatalf("setting object field: %v", err) + } + input["x"].([]any)[0] = 9 + if got := session.Get("obj").(map[string]any)["x"].([]any)[0]; got != 1.0 { + t.Fatalf("stored object changed after mutating input: got %v, want 1", got) + } + output := session.Get("obj").(map[string]any) + output["x"].([]any)[0] = 8 + if got := session.Get("obj").(map[string]any)["x"].([]any)[0]; got != 1.0 { + t.Fatalf("stored object changed after mutating returned value: got %v, want 1", got) + } +} + +func TestGetMissingKey(t *testing.T) { + _, _, manager := newTest(t, time.Minute) + session, _ := createSession(t, manager) + if got := session.Get("missing"); got != nil { + t.Fatalf("missing field = %v, want nil", got) + } +} + +func TestSessionMutationRefreshesTTL(t *testing.T) { + const configuredTTL = 60 * time.Second + miniRedis, client, manager := newTest(t, configuredTTL) + session, sessionID := createSession(t, manager) + operations := []struct { + name string + call func() error + wantErr error + }{ + { + name: "set", + call: func() error { return session.Set("x", 1) }, + wantErr: nil, + }, + { + name: "delete", + call: func() error { return session.Delete("x") }, + wantErr: nil, + }, + { + name: "update TTL", + call: session.UpdateTTL, + wantErr: nil, + }, + } + for _, operation := range operations { + t.Run(operation.name, func(t *testing.T) { + miniRedis.FastForward(20 * time.Second) + gotErr := operation.call() + if !errors.Is(gotErr, operation.wantErr) { + t.Fatalf("%s operation error = %v, want %v", operation.name, gotErr, operation.wantErr) + } + remainingTTL, err := client.PTTL(context.Background(), sessionKeyPrefix+sessionID).Result() + if err != nil { + t.Fatalf("reading TTL after %s: %v", operation.name, err) + } + wantMinTTL, wantMaxTTL := configuredTTL-time.Second, configuredTTL + if remainingTTL < wantMinTTL || remainingTTL > wantMaxTTL { + t.Fatalf("TTL after %s = %v, want within one second of %v", operation.name, remainingTTL, configuredTTL) + } + }) + } +} + +func TestSessionScriptErrorPreservesSession(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + if err := session.Set("x", "ok"); err != nil { + t.Fatalf("seeding ready session: %v", err) + } + sentinelError := errors.New("boom") + client.AddHook(&redisTestHook{commandErrs: map[string]error{"evalsha": sentinelError, "eval": sentinelError}}) + + gotErr := session.Set("y", 1) + if !errors.Is(gotErr, sentinelError) { + t.Fatalf("failed mutation error = %v, want sentinel error", gotErr) + } + if got := session.Get("x"); got != "ok" { + t.Fatalf("ready session lost existing value after script error: got %v, want %q", got, "ok") + } + if _, found := manager.cachedSession(sessionID); !found { + t.Fatal("ready session was removed from cache after script error") + } +} + +func TestFirstLoadEvictsPlaceholder(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + sentinelError := errors.New("boom") + client.AddHook(&redisTestHook{commandErrs: map[string]error{"evalsha": sentinelError, "eval": sentinelError}}) + _, _, gotErr := getSessionRequest(t, manager, "missing") + if !errors.Is(gotErr, sentinelError) { + t.Fatalf("first load error = %v, want sentinel error", gotErr) + } + if len(manager.cache) != 0 { + t.Fatalf("failed first load left %d cache entries, want none", len(manager.cache)) + } +} + +func TestInvalidHashDataPreservesForeignData(t *testing.T) { + miniRedis, client, manager := newTest(t, time.Minute) + invalidHashCases := []struct { + name string + seed func() + }{ + { + name: "unsupported schema", + seed: func() { miniRedis.HSet(sessionKeyPrefix+"x", metadataField, "2") }, + }, + { + name: "unknown hash field", + seed: func() { miniRedis.HSet(sessionKeyPrefix+"x", metadataField, "1", "junk", "1") }, + }, + { + name: "malformed JSON", + seed: func() { + miniRedis.HSet(sessionKeyPrefix+"x", metadataField, "1", encodeUserField("x"), "not-json") + }, + }, + } + for _, testCase := range invalidHashCases { + t.Run(testCase.name, func(t *testing.T) { + miniRedis.FlushAll() + manager = NewManager(client, "sid", time.Minute, false) + testCase.seed() + _, _, err := getSessionRequest(t, manager, "x") + if !errors.Is(err, ErrInvalidSessionData) { + t.Fatalf("GetSession error = %v, want ErrInvalidSessionData", err) + } + }) + } + + const foreignTTL = 30 * time.Second + const foreignString = `{"x":1}` + foreignDataCases := []struct { + name string + seed func() error + wantType string + wantValues []string + }{ + { + name: "foreign list", + seed: func() error { + if err := client.RPush(context.Background(), sessionKeyPrefix+"x", "bad", "data").Err(); err != nil { + return err + } + return client.Expire(context.Background(), sessionKeyPrefix+"x", foreignTTL).Err() + }, + wantType: "list", + wantValues: []string{"bad", "data"}, + }, + { + name: "foreign string", + seed: func() error { + return client.Set(context.Background(), sessionKeyPrefix+"x", foreignString, foreignTTL).Err() + }, + wantType: "string", + wantValues: []string{foreignString}, + }, + } + for _, testCase := range foreignDataCases { + t.Run(testCase.name, func(t *testing.T) { + miniRedis.FlushAll() + manager = NewManager(client, "sid", time.Minute, false) + if err := testCase.seed(); err != nil { + t.Fatalf("seeding foreign %s: %v", testCase.name, err) + } + + capture := func() (string, []string, time.Duration) { + key := sessionKeyPrefix + "x" + valueType, err := client.Type(context.Background(), key).Result() + if err != nil { + t.Fatalf("reading foreign %s type: %v", testCase.name, err) + } + var values []string + switch valueType { + case "list": + values, err = client.LRange(context.Background(), key, 0, -1).Result() + case "string": + var value string + value, err = client.Get(context.Background(), key).Result() + values = []string{value} + default: + t.Fatalf("foreign %s type = %q, want list or string", testCase.name, valueType) + } + if err != nil { + t.Fatalf("reading foreign %s value: %v", testCase.name, err) + } + ttl, err := client.PTTL(context.Background(), key).Result() + if err != nil { + t.Fatalf("reading foreign %s TTL: %v", testCase.name, err) + } + return valueType, values, ttl + } + + beforeType, beforeValues, beforeTTL := capture() + if beforeType != testCase.wantType { + t.Fatalf("foreign %s type before GetSession = %q, want %q", testCase.name, beforeType, testCase.wantType) + } + if !reflect.DeepEqual(beforeValues, testCase.wantValues) { + t.Fatalf("foreign %s value before GetSession = %q, want %q", testCase.name, beforeValues, testCase.wantValues) + } + if beforeTTL <= 0 { + t.Fatalf("foreign %s TTL before GetSession = %v, want positive TTL", testCase.name, beforeTTL) + } + + _, recorder, err := getSessionRequest(t, manager, "x") + if !errors.Is(err, ErrInvalidSessionData) { + t.Fatalf("GetSession error = %v, want ErrInvalidSessionData", err) + } + if cookies := recorder.Result().Cookies(); len(cookies) != 0 { + t.Fatalf("invalid %s session emitted %d replacement cookies, want none", testCase.name, len(cookies)) + } + + afterType, afterValues, afterTTL := capture() + if afterType != beforeType { + t.Fatalf("foreign %s type changed from %q to %q", testCase.name, beforeType, afterType) + } + if !reflect.DeepEqual(afterValues, beforeValues) { + t.Fatalf("foreign %s value changed from %q to %q", testCase.name, beforeValues, afterValues) + } + if afterTTL < beforeTTL-time.Second || afterTTL > beforeTTL+time.Second { + t.Fatalf("foreign %s TTL changed from %v to %v", testCase.name, beforeTTL, afterTTL) + } + }) + } +} + +func TestMalformedSessionReadPreservesTTL(t *testing.T) { + const managerTTL = time.Minute + const shortTTL = 10 * time.Second + miniRedis, client, manager := newTest(t, managerTTL) + sessionID := "malformed-ttl" + sessionKey := sessionKeyPrefix + sessionID + miniRedis.HSet(sessionKey, metadataField, schemaVersion, encodeUserField("user"), "not-json") + if ok, err := client.Expire(context.Background(), sessionKey, shortTTL).Result(); err != nil || !ok { + t.Fatalf("seeding malformed-session TTL: ok=%v err=%v", ok, err) + } + + _, _, err := getSessionRequest(t, manager, sessionID) + if !errors.Is(err, ErrInvalidSessionData) { + t.Fatalf("GetSession error = %v, want ErrInvalidSessionData", err) + } + remainingTTL, err := client.PTTL(context.Background(), sessionKey).Result() + if err != nil { + t.Fatalf("reading malformed-session TTL: %v", err) + } + if remainingTTL <= 0 || remainingTTL >= managerTTL/2 { + t.Fatalf("malformed-session TTL = %v, want positive TTL well below manager TTL", remainingTTL) + } +} + +func TestDestroyCanBeCalledTwice(t *testing.T) { + miniRedis, _, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + if err := session.Destroy(); err != nil { + t.Fatalf("first Destroy call: %v", err) + } + if err := session.Destroy(); err != nil { + t.Fatalf("second Destroy call: %v", err) + } + if session.ID() != sessionID { + t.Fatalf("destroyed session ID = %q, want original %q", session.ID(), sessionID) + } + if session.Get("x") != nil { + t.Fatal("destroyed session returned a stored value") + } + if _, found := manager.cachedSession(sessionID); found { + t.Fatal("destroyed session remained cached") + } + if miniRedis.Exists(sessionKeyPrefix + sessionID) { + t.Fatal("destroyed session Redis key still exists") + } +} + +func TestExpiredSessionEvictsIdleEntry(t *testing.T) { + miniRedis, _, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + miniRedis.Del(sessionKeyPrefix + sessionID) + expireSessionForTest(session) + if _, found := manager.cachedSession(sessionID); found { + t.Fatal("expired idle session remained cached after Redis key disappeared") + } +} + +func TestFirstLoadOperationalErrorReachesCaller(t *testing.T) { + miniRedis, client, manager := newTest(t, time.Minute) + sessionID := "operational-error" + miniRedis.HSet(sessionKeyPrefix+sessionID, metadataField, schemaVersion) + sentinelError := errors.New("sentinel redis failure") + client.AddHook(&redisTestHook{commandErrs: map[string]error{"evalsha": sentinelError, "eval": sentinelError}}) + + placeholder := seedCachedSession(manager, &Session{sid: sessionID, manager: manager}) + if err := placeholder.redisReadAndRefresh(context.Background()); !errors.Is(err, sentinelError) { + t.Fatalf("first load error = %v, want sentinel error", err) + } + if !placeholder.isClosed() { + t.Fatal("failed first-load placeholder was not marked closed") + } + if cached, found := manager.cachedSession(sessionID); !found || cached != placeholder { + t.Fatal("failed first-load placeholder was not retained for the waiting caller") + } + + session, recorder, err := getSessionRequest(t, manager, sessionID) + if session != nil { + t.Fatalf("waiting caller returned session %v, want nil", session) + } + if !errors.Is(err, sentinelError) { + t.Fatalf("waiting caller error = %v, want sentinel error", err) + } + if cookies := recorder.Result().Cookies(); len(cookies) != 0 { + t.Fatalf("waiting caller emitted %d replacement cookies, want none", len(cookies)) + } + if _, found := manager.cachedSession(sessionID); found { + t.Fatal("failed first-load placeholder remained cached after waiting caller cleanup") + } +} + +func TestFirstLoadInvalidDataReachesCaller(t *testing.T) { + miniRedis, _, manager := newTest(t, time.Minute) + sessionID := "invalid-data" + sessionKey := sessionKeyPrefix + sessionID + miniRedis.HSet(sessionKey, metadataField, schemaVersion, "foreign", "value") + + placeholder := seedCachedSession(manager, &Session{sid: sessionID, manager: manager}) + if err := placeholder.redisReadAndRefresh(context.Background()); !errors.Is(err, ErrInvalidSessionData) { + t.Fatalf("first load error = %v, want ErrInvalidSessionData", err) + } + if !placeholder.isClosed() { + t.Fatal("invalid first-load placeholder was not marked closed") + } + if cached, found := manager.cachedSession(sessionID); !found || cached != placeholder { + t.Fatal("invalid first-load placeholder was not retained for the waiting caller") + } + + session, recorder, err := getSessionRequest(t, manager, sessionID) + if session != nil { + t.Fatalf("waiting caller returned session %v, want nil", session) + } + if !errors.Is(err, ErrInvalidSessionData) { + t.Fatalf("waiting caller error = %v, want ErrInvalidSessionData", err) + } + if cookies := recorder.Result().Cookies(); len(cookies) != 0 { + t.Fatalf("waiting caller emitted %d replacement cookies, want none", len(cookies)) + } + if _, found := manager.cachedSession(sessionID); found { + t.Fatal("invalid first-load placeholder remained cached after waiting caller cleanup") + } + if !miniRedis.Exists(sessionKey) || miniRedis.HGet(sessionKey, "foreign") != "value" { + t.Fatal("invalid foreign Redis data was deleted or changed") + } +} + +func TestNewManagerCookieName(t *testing.T) { + testCases := []struct { + name string + cookieName string + shouldBeValid bool + }{ + { + name: "empty name", + cookieName: "", + shouldBeValid: false, + }, + { + name: "invalid name", + cookieName: "bad name", + shouldBeValid: false, + }, + { + name: "valid name", + cookieName: "sid", + shouldBeValid: true, + }, + } + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + if testCase.shouldBeValid { + if manager := NewManager(nil, testCase.cookieName, time.Minute, false); manager == nil { + t.Fatal("NewManager returned nil for valid cookie name") + } + return + } + defer func() { + if recovered := recover(); recovered != "session: invalid cookie name" { + t.Fatalf("panic = %v, want exact invalid-cookie panic", recovered) + } + }() + NewManager(nil, testCase.cookieName, time.Minute, false) + }) + } +} + +func TestNewManagerRejectsShortTTL(t *testing.T) { + defer func() { + if recovered := recover(); recovered != "session: ttl must be at least 1 second" { + t.Fatalf("panic = %v, want minimum-TTL validation panic", recovered) + } + }() + NewManager(nil, "sid", time.Second-time.Nanosecond, false) +} + +func TestExpireKeepsSessionAfterRefresh(t *testing.T) { + _, client, firstManager := newTest(t, time.Minute) + secondManager := NewManager(client, "sid", time.Minute, false) + firstSession, sessionID := createSession(t, firstManager) + secondSession, _, err := getSessionRequest(t, secondManager, sessionID) + if err != nil { + t.Fatalf("loading session through second manager: %v", err) + } + if err := secondSession.UpdateTTL(); err != nil { + t.Fatalf("refreshing TTL through second manager: %v", err) + } + + expireSessionForTest(firstSession) + if _, found := firstManager.cachedSession(sessionID); !found { + t.Fatal("local expiry evicted a session refreshed by another manager") + } +} + +func TestParseReadResult(t *testing.T) { + var command redis.Cmd + command.SetVal([]any{int64(1), metadataField, "1", encodeUserField("n"), "2"}) + + schema, values, err := parseReadResult(&command) + if err != nil { + t.Fatalf("parseReadResult returned an error: %v", err) + } + if schema != 1 { + t.Fatalf("parsed schema version = %d, want 1", schema) + } + if string(values["n"]) != "2" { + t.Fatalf("parsed user field n = %q, want %q", values["n"], "2") + } +} + +func TestExpiredSessionMutations(t *testing.T) { + testCases := []struct { + name string + mutate func(*Session) error + wantErr error + }{ + { + name: "set", + mutate: func(session *Session) error { return session.Set("x", 1) }, + wantErr: ErrSessionClosed, + }, + { + name: "delete", + mutate: func(session *Session) error { return session.Delete("x") }, + wantErr: ErrSessionClosed, + }, + { + name: "update TTL", + mutate: func(session *Session) error { return session.UpdateTTL() }, + wantErr: redis.Nil, + }, + } + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + miniRedis, _, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + miniRedis.Del(sessionKeyPrefix + sessionID) + + gotErr := testCase.mutate(session) + if !errors.Is(gotErr, testCase.wantErr) { + t.Fatalf("expired %s error = %v, want %v", testCase.name, gotErr, testCase.wantErr) + } + if miniRedis.Exists(sessionKeyPrefix + sessionID) { + t.Fatal("expired session mutation resurrected the Redis key") + } + if _, found := manager.cachedSession(sessionID); found { + t.Fatal("closed session remained cached after expired mutation") + } + }) + } +} + +func TestHashFieldEncodingRoundTrips(t *testing.T) { + _, _, manager := newTest(t, time.Minute) + session, _ := createSession(t, manager) + fieldNames := []string{"", "世界/ punctuation!", "_meta", "v:abc"} + for _, fieldName := range fieldNames { + if err := session.Set(fieldName, 7); err != nil { + t.Fatalf("setting field %q: %v", fieldName, err) + } + value, ok := session.Get(fieldName).(float64) + if !ok || value != 7 { + t.Fatalf("field %q decoded as %#v, want float64(7)", fieldName, session.Get(fieldName)) + } + } +} + +func TestSameHashFieldUsesLastWrite(t *testing.T) { + _, client, firstManager := newTest(t, time.Minute) + secondManager := NewManager(client, "sid", time.Minute, false) + _, sessionID := createSession(t, firstManager) + firstSession, _, err := getSessionRequest(t, firstManager, sessionID) + if err != nil { + t.Fatalf("loading first session: %v", err) + } + secondSession, _, err := getSessionRequest(t, secondManager, sessionID) + if err != nil { + t.Fatalf("loading second session: %v", err) + } + if err := firstSession.Set("same", "first"); err != nil { + t.Fatalf("first write: %v", err) + } + if err := secondSession.Set("same", "second"); err != nil { + t.Fatalf("second write: %v", err) + } + if got := secondSession.Get("same"); got != "second" { + t.Fatalf("same-field value = %v, want %q after last successful write", got, "second") + } +} + +func TestExpireRetriesAfterPTTLFailure(t *testing.T) { + _, client, manager := newTest(t, time.Minute) + session, sessionID := createSession(t, manager) + client.AddHook(&redisTestHook{commandErrs: map[string]error{"pttl": errors.New("temporary")}}) + expireSessionForTest(session) + if _, found := manager.cachedSession(sessionID); !found { + t.Fatal("PTTL failure evicted the session instead of retaining it") + } + session.mutex.RLock() + retryScheduled := session.expiresAt.After(time.Now()) + session.mutex.RUnlock() + if !retryScheduled { + t.Fatal("PTTL failure did not schedule a future expiry retry") + } +} + +func ExampleManager_CreateSession() { + miniRedis, err := miniredis.Run() + if err != nil { + return + } + defer miniRedis.Close() + client := redis.NewClient(&redis.Options{Addr: miniRedis.Addr()}) + defer client.Close() + manager := NewManager(client, "session_id", 96*time.Hour, true) + request := httptest.NewRequest(http.MethodGet, "/", nil) + recorder := httptest.NewRecorder() + _, _ = manager.CreateSession(recorder, request) +} + +func ExampleManager_GetSession() { + miniRedis, err := miniredis.Run() + if err != nil { + return + } + defer miniRedis.Close() + client := redis.NewClient(&redis.Options{Addr: miniRedis.Addr()}) + defer client.Close() + manager := NewManager(client, "session_id", 96*time.Hour, true) + createRequest := httptest.NewRequest(http.MethodGet, "/", nil) + createRecorder := httptest.NewRecorder() + _, _ = manager.CreateSession(createRecorder, createRequest) + request := httptest.NewRequest(http.MethodGet, "/", nil) + for _, cookie := range createRecorder.Result().Cookies() { + request.AddCookie(cookie) + } + recorder := httptest.NewRecorder() + _, _ = manager.GetSession(recorder, request) +}