From 162e67a935e61dd76966a27cb599f381c730a565 Mon Sep 17 00:00:00 2001 From: Patrick East Date: Tue, 9 Apr 2019 16:31:18 -0700 Subject: [PATCH] Add unit tests for http server shutdown timeout To make this work it adds in a new shim interface for the server to interact with the `http.Server`(s). We mock one out and can then have it return errors on `Shutdown`. Signed-off-by: Patrick East --- server/server.go | 64 +++++++++++++++++++++++++---------- server/server_test.go | 78 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 124 insertions(+), 18 deletions(-) diff --git a/server/server.go b/server/server.go index 1efcd7d078..d8fbf08f3f 100644 --- a/server/server.go +++ b/server/server.go @@ -104,7 +104,7 @@ type Server struct { errLimit int pprofEnabled bool runtime *ast.Term - httpServers []*http.Server + httpListeners []httpListener } // Loop will contain all the calls from the server that we'll be listening on. @@ -172,14 +172,14 @@ func (s *Server) Init(ctx context.Context) (*Server, error) { // by the context an error will be returned. func (s *Server) Shutdown(ctx context.Context) error { errChan := make(chan error) - for _, srvr := range s.httpServers { - go func(s *http.Server) { + for _, srvr := range s.httpListeners { + go func(s httpListener) { errChan <- s.Shutdown(ctx) }(srvr) } // wait until each server has finished shutting down var errorList []error - for i := 0; i < len(s.httpServers); i++ { + for i := 0; i < len(s.httpListeners); i++ { err := <-errChan if err != nil { errorList = append(errorList, err) @@ -307,21 +307,21 @@ func (s *Server) Listeners() ([]Loop, error) { return nil, err } var loop Loop - var httpServer *http.Server + var listener httpListener switch parsedURL.Scheme { case "unix": - loop, httpServer, err = s.getListenerForUNIXSocket(parsedURL) + loop, listener, err = s.getListenerForUNIXSocket(parsedURL) case "http": - loop, httpServer, err = s.getListenerForHTTPServer(parsedURL) + loop, listener, err = s.getListenerForHTTPServer(parsedURL) case "https": - loop, httpServer, err = s.getListenerForHTTPSServer(parsedURL) + loop, listener, err = s.getListenerForHTTPSServer(parsedURL) default: err = fmt.Errorf("invalid url scheme %q", parsedURL.Scheme) } if err != nil { return nil, err } - s.httpServers = append(s.httpServers, httpServer) + s.httpListeners = append(s.httpListeners, listener) loops = append(loops, loop) } @@ -330,28 +330,56 @@ func (s *Server) Listeners() ([]Loop, error) { if err != nil { return nil, err } - var httpServer *http.Server - loop, httpServer, err := s.getListenerForHTTPServer(parsedURL) + loop, httpListener, err := s.getListenerForHTTPServer(parsedURL) if err != nil { return nil, err } - s.httpServers = append(s.httpServers, httpServer) + s.httpListeners = append(s.httpListeners, httpListener) loops = append(loops, loop) } return loops, nil } -func (s *Server) getListenerForHTTPServer(u *url.URL) (Loop, *http.Server, error) { +type httpListener interface { + ListenAndServe() error + ListenAndServeTLS(certFile, keyFile string) error + Shutdown(ctx context.Context) error +} + +// baseHTTPListener is just a wrapper around http.Server +type baseHTTPListener struct { + s *http.Server +} + +var _ httpListener = (*baseHTTPListener)(nil) + +func newHTTPListener(srvr *http.Server) httpListener { + return &baseHTTPListener{srvr} +} + +func (l *baseHTTPListener) ListenAndServe() error { + return l.s.ListenAndServe() +} + +func (l *baseHTTPListener) ListenAndServeTLS(certFile, keyFile string) error { + return l.s.ListenAndServeTLS(certFile, keyFile) +} + +func (l *baseHTTPListener) Shutdown(ctx context.Context) error { + return l.s.Shutdown(ctx) +} + +func (s *Server) getListenerForHTTPServer(u *url.URL) (Loop, httpListener, error) { httpServer := http.Server{ Addr: u.Host, Handler: s.Handler, } - return httpServer.ListenAndServe, &httpServer, nil + return httpServer.ListenAndServe, newHTTPListener(&httpServer), nil } -func (s *Server) getListenerForHTTPSServer(u *url.URL) (Loop, *http.Server, error) { +func (s *Server) getListenerForHTTPSServer(u *url.URL) (Loop, httpListener, error) { if s.cert == nil { return nil, nil, fmt.Errorf("TLS certificate required but not supplied") @@ -371,10 +399,10 @@ func (s *Server) getListenerForHTTPSServer(u *url.URL) (Loop, *http.Server, erro httpsLoop := func() error { return httpsServer.ListenAndServeTLS("", "") } - return httpsLoop, &httpsServer, nil + return httpsLoop, newHTTPListener(&httpsServer), nil } -func (s *Server) getListenerForUNIXSocket(u *url.URL) (Loop, *http.Server, error) { +func (s *Server) getListenerForUNIXSocket(u *url.URL) (Loop, httpListener, error) { socketPath := u.Host + u.Path // Remove domain socket file in case it already exists. @@ -387,7 +415,7 @@ func (s *Server) getListenerForUNIXSocket(u *url.URL) (Loop, *http.Server, error } domainSocketLoop := func() error { return domainSocketServer.Serve(unixListener) } - return domainSocketLoop, &domainSocketServer, nil + return domainSocketLoop, newHTTPListener(&domainSocketServer), nil } func (s *Server) initRouter() { diff --git a/server/server_test.go b/server/server_test.go index c4d491b552..b136f114c8 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -3358,3 +3358,81 @@ func TestShutdown(t *testing.T) { t.Errorf("unexpected error shutting down server: %s", err.Error()) } } + +func TestShutdownError(t *testing.T) { + f := newFixture(t) + + errMsg := "failed to shutdown" + + // Add a mock httpListener to the server + m := &mockHTTPListener{ + ShutdownHook: func() error { + return errors.New(errMsg) + }, + } + f.server.httpListeners = []httpListener{m} + + ctx, cancel := context.WithTimeout(context.Background(), time.Duration(5)*time.Second) + defer cancel() + err := f.server.Shutdown(ctx) + if err == nil { + t.Error("expected an error shutting down server but err==nil") + } else if !strings.Contains(err.Error(), errMsg) { + t.Errorf("unexpected error shutting down server: %s", err.Error()) + } +} + +func TestShutdownMultipleErrors(t *testing.T) { + f := newFixture(t) + + shutdownErrs := []error{errors.New("err1"), nil, errors.New("err3")} + + // Add mock httpListeners to the server + for _, err := range shutdownErrs { + m := &mockHTTPListener{} + if err != nil { + retVal := errors.New(err.Error()) + m.ShutdownHook = func() error { + return retVal + } + } + f.server.httpListeners = append(f.server.httpListeners, m) + } + + ctx, cancel := context.WithTimeout(context.Background(), time.Duration(5)*time.Second) + defer cancel() + err := f.server.Shutdown(ctx) + if err == nil { + t.Fatal("expected an error shutting down server but err==nil") + } + + for _, expectedErr := range shutdownErrs { + if expectedErr != nil && !strings.Contains(err.Error(), expectedErr.Error()) { + t.Errorf("expected error message to contain '%s', full message: '%s'", expectedErr.Error(), err.Error()) + } + } +} + +type listenerHook func() error + +type mockHTTPListener struct { + ShutdownHook listenerHook +} + +var _ httpListener = (*mockHTTPListener)(nil) + +func (m mockHTTPListener) ListenAndServe() error { + return errors.New("not implemented") +} + +func (m mockHTTPListener) ListenAndServeTLS(certFile, keyFile string) error { + return errors.New("not implemented") +} + +func (m mockHTTPListener) Shutdown(ctx context.Context) error { + var err error + if m.ShutdownHook != nil { + err = m.ShutdownHook() + } + return err +}