Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions sessions.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"encoding/gob"
"fmt"
"net/http"
"sync"
"time"
)

Expand Down Expand Up @@ -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 {
Expand All @@ -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
}
Expand All @@ -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 {
Expand All @@ -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
Expand Down
26 changes: 26 additions & 0 deletions sessions_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
)

Expand Down Expand Up @@ -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()
Expand Down
Loading