diff --git a/sessions.go b/sessions.go index c052b28..e08bd70 100644 --- a/sessions.go +++ b/sessions.go @@ -137,6 +137,9 @@ 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 { + 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..fe8849f 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" @@ -214,6 +215,56 @@ func TestCookieStoreMapPanic(t *testing.T) { } } +type errorStore struct{} + +func (errorStore) Get(r *http.Request, name string) (*Session, error) { + return GetRegistry(r).Get(errorStore{}, name) +} + +func (errorStore) New(r *http.Request, name string) (*Session, error) { + return nil, errors.New("store error") +} + +func (errorStore) Save(r *http.Request, w http.ResponseWriter, s *Session) error { + return nil +} + +func TestRegistryGetStoreNewError(t *testing.T) { + defer func() { + if r := recover(); r != nil { + t.Fatalf("Registry.Get panicked: %v", r) + } + }() + + req, err := http.NewRequest("GET", "http://www.example.com", nil) + if err != nil { + t.Fatal("failed to create request", err) + } + + session, err := GetRegistry(req).Get(errorStore{}, "test-session") + if session != nil { + t.Fatalf("expected nil session, got %#v", session) + } + if err == nil || err.Error() != "store error" { + t.Fatalf("expected 'store error', got %v", err) + } + + // Verify Save on registry does not panic when a session failed to initialize. + w := httptest.NewRecorder() + if err := Save(req, w); err != nil { + t.Fatalf("unexpected error saving registry: %v", err) + } + + // Verify calling Get again returns the error rather than panicking. + session2, err2 := GetRegistry(req).Get(errorStore{}, "test-session") + if session2 != nil { + t.Fatalf("expected nil session on second call, got %#v", session2) + } + if err2 == nil || err2.Error() != "store error" { + t.Fatalf("expected 'store error' on second call, got %v", err2) + } +} + func init() { gob.Register(FlashMessage{}) }