Files
2026-09-06 22:32:59 +03:00

1138 lines
37 KiB
Go

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