diff --git a/sessions.go b/sessions.go index c052b28..2c1fec7 100644 --- a/sessions.go +++ b/sessions.go @@ -137,6 +137,12 @@ 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 { + // Some stores return (nil, err) when initialization fails; + // surface the error rather than panicking on session.name. + s.sessions[name] = sessionInfo{s: nil, e: err} + return nil, err + } session.name = name s.sessions[name] = sessionInfo{s: session, e: err} } diff --git a/sessions_test.go b/sessions_test.go index 9476c22..9245e5e 100644 --- a/sessions_test.go +++ b/sessions_test.go @@ -214,6 +214,49 @@ func TestCookieStoreMapPanic(t *testing.T) { } } +// nilSessionStore returns (nil, err) from New, which the registry used +// to dereference unconditionally and panic on. +type nilSessionStore struct{} + +func (nilSessionStore) Get(_ *http.Request, _ string) (*Session, error) { + return nil, errSentinel +} + +func (nilSessionStore) New(_ *http.Request, _ string) (*Session, error) { + return nil, errSentinel +} + +func (nilSessionStore) Save(_ *http.Request, _ http.ResponseWriter, _ *Session) error { + return errSentinel +} + +var errSentinel = stringError("store unavailable") + +type stringError string + +func (e stringError) Error() string { return string(e) } + +func TestRegistryGetReturnsErrorWhenStoreReturnsNilSession(t *testing.T) { + defer func() { + if r := recover(); r != nil { + t.Fatalf("Registry.Get panicked when store returned nil: %v", r) + } + }() + + req, err := http.NewRequest("GET", "http://www.example.com", nil) + if err != nil { + t.Fatal(err) + } + reg := GetRegistry(req) + sess, err := reg.Get(nilSessionStore{}, "name") + if err == nil { + t.Fatal("expected error from Registry.Get, got nil") + } + if sess != nil { + t.Fatalf("expected nil session, got %+v", sess) + } +} + func init() { gob.Register(FlashMessage{}) }