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) }