From 4768d38ca4b6df1807e6a0959513d7ab6d979fc7 Mon Sep 17 00:00:00 2001 From: AshSgDe29071999 Date: Mon, 31 Aug 2026 14:57:09 +0530 Subject: [PATCH] Serialize session registry access per request Concurrent Get calls on the same request raced on Registry.sessions and on replacing the request context. GraphQL subscriptions hitting sessions.Get from several goroutines could fatal with concurrent map writes. See #287 --- sessions.go | 10 ++++++++++ sessions_test.go | 26 ++++++++++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/sessions.go b/sessions.go index c052b28..fd1c583 100644 --- a/sessions.go +++ b/sessions.go @@ -9,6 +9,7 @@ import ( "encoding/gob" "fmt" "net/http" + "sync" "time" ) @@ -105,8 +106,12 @@ type contextKey int // registryKey is the key used to store the registry in the context. const registryKey contextKey = 0 +var registryMu sync.Mutex + // GetRegistry returns a registry instance for the current request. func GetRegistry(r *http.Request) *Registry { + registryMu.Lock() + defer registryMu.Unlock() var ctx = r.Context() registry := ctx.Value(registryKey) if registry != nil { @@ -122,6 +127,7 @@ func GetRegistry(r *http.Request) *Registry { // Registry stores sessions used during a request. type Registry struct { + mu sync.Mutex request *http.Request sessions map[string]sessionInfo } @@ -133,6 +139,8 @@ func (s *Registry) Get(store Store, name string) (session *Session, err error) { if !isCookieNameValid(name) { return nil, fmt.Errorf("sessions: invalid character in cookie name: %s", name) } + s.mu.Lock() + defer s.mu.Unlock() if info, ok := s.sessions[name]; ok { session, err = info.s, info.e } else { @@ -146,6 +154,8 @@ func (s *Registry) Get(store Store, name string) (session *Session, err error) { // Save saves all sessions registered for the current request. func (s *Registry) Save(w http.ResponseWriter) error { + s.mu.Lock() + defer s.mu.Unlock() var errMulti MultiError for name, info := range s.sessions { session := info.s diff --git a/sessions_test.go b/sessions_test.go index 9476c22..ff61fde 100644 --- a/sessions_test.go +++ b/sessions_test.go @@ -10,6 +10,7 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" ) @@ -189,6 +190,31 @@ func TestFlashes(t *testing.T) { } } +func TestRegistryGetConcurrent(t *testing.T) { + store := NewCookieStore([]byte("aaa0defe5d2839cbc46fc4f080cd7adc")) + req, err := http.NewRequest("GET", "http://www.example.com", nil) + if err != nil { + t.Fatal(err) + } + + var wg sync.WaitGroup + errCh := make(chan error, 32) + for range 32 { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := store.Get(req, "sess"); err != nil { + errCh <- err + } + }() + } + wg.Wait() + close(errCh) + for err := range errCh { + t.Fatal(err) + } +} + func TestCookieStoreMapPanic(t *testing.T) { defer func() { err := recover()