diff --git a/sessions.go b/sessions.go index c052b28..aed2359 100644 --- a/sessions.go +++ b/sessions.go @@ -137,9 +137,18 @@ func (s *Registry) Get(store Store, name string) (session *Session, err error) { session, err = info.s, info.e } else { session, err = store.New(s.request, name) + if session == nil { + if err == nil { + err = fmt.Errorf("sessions: store.New returned a nil session for %s", name) + } + return nil, err + } session.name = name s.sessions[name] = sessionInfo{s: session, e: err} } + if session == nil { + return nil, err + } session.store = store return } diff --git a/sessions_test.go b/sessions_test.go index 9476c22..299e08d 100644 --- a/sessions_test.go +++ b/sessions_test.go @@ -7,6 +7,7 @@ package sessions import ( "bytes" "encoding/gob" + "errors" "net/http" "net/http/httptest" "strings" @@ -189,6 +190,34 @@ func TestFlashes(t *testing.T) { } } +type errNilStore struct{} + +func (errNilStore) Get(r *http.Request, name string) (*Session, error) { + return GetRegistry(r).Get(errNilStore{}, name) +} + +func (errNilStore) New(*http.Request, string) (*Session, error) { + return nil, errors.New("store boom") +} + +func (errNilStore) Save(*http.Request, http.ResponseWriter, *Session) error { + return nil +} + +func TestRegistryGetNilSessionFromStore(t *testing.T) { + req, err := http.NewRequest("GET", "http://localhost/", nil) + if err != nil { + t.Fatal(err) + } + session, err := GetRegistry(req).Get(errNilStore{}, "session-key") + if session != nil { + t.Fatalf("expected nil session, got %#v", session) + } + if err == nil || err.Error() != "store boom" { + t.Fatalf("expected store boom, got %v", err) + } +} + func TestCookieStoreMapPanic(t *testing.T) { defer func() { err := recover()