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 <east.patrick@gmail.com>
This commit is contained in:
Patrick East
2019-04-09 16:31:18 -07:00
committed by Torin Sandall
parent 713f961bfe
commit 162e67a935
2 changed files with 124 additions and 18 deletions
+46 -18
View File
@@ -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() {
+78
View File
@@ -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
}