From 414e004bbdc5b1c3c2b3933356b241d47719ed51 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Thu, 9 Nov 2017 18:43:44 +0300 Subject: [PATCH 1/8] ldap_server_shutdown | add Server.Shutdown method --- ldap/server.go | 56 ++++++++++++++++++++++++++++++++++++++-- ldap/server_test.go | 62 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 116 insertions(+), 2 deletions(-) create mode 100644 ldap/server_test.go diff --git a/ldap/server.go b/ldap/server.go index df17675..9709ede 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -10,6 +10,7 @@ import ( "net" "os" "strings" + "sync" ) func NewResponsePacket(msgID int) *Packet { @@ -80,6 +81,10 @@ type Server struct { RootDSE map[string][]string tlsConfig *tls.Config + mu sync.Mutex + listeners []net.Listener + stopC chan struct{} + factory listenerFactory } type srvClient struct { @@ -89,6 +94,20 @@ type srvClient struct { ctx Context } +type listenerFactory interface { + newListener(network, address string) (net.Listener, error) + newTLSListener(network, addr string, config *tls.Config) (net.Listener, error) +} + +type listenerFactoryImpl struct{} + +func (_ *listenerFactoryImpl) newListener(network, address string) (net.Listener, error) { + return net.Listen(network, address) +} +func (_ *listenerFactoryImpl) newTLSListener(network, addr string, config *tls.Config) (net.Listener, error) { + return tls.Listen(network, addr, config) +} + func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) { // Copy the default RootDSE sf := make(map[string][]string, len(RootDSE)) @@ -106,9 +125,26 @@ func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) { Backend: be, RootDSE: sf, tlsConfig: tlsConfig, + factory: &listenerFactoryImpl{}, + stopC: make(chan struct{}), }, nil } +func (srv *Server) Shutdown() error { + var err error + + close(srv.stopC) + srv.mu.Lock() + for _, listener := range srv.listeners { + if e := listener.Close(); err != nil { + err = e + } + } + srv.mu.Unlock() + + return err +} + func (srv *Server) ServeTLS(network, addr string, tlsConfig *tls.Config) error { if tlsConfig == nil { tlsConfig = srv.tlsConfig @@ -116,7 +152,7 @@ func (srv *Server) ServeTLS(network, addr string, tlsConfig *tls.Config) error { if tlsConfig == nil { return errors.New("ldap: no TLS config") } - ln, err := tls.Listen(network, addr, tlsConfig) + ln, err := srv.factory.newTLSListener(network, addr, tlsConfig) if err != nil { return err } @@ -124,7 +160,7 @@ func (srv *Server) ServeTLS(network, addr string, tlsConfig *tls.Config) error { } func (srv *Server) Serve(network, addr string) error { - ln, err := net.Listen(network, addr) + ln, err := srv.factory.newListener(network, addr) if err != nil { return err } @@ -132,7 +168,23 @@ func (srv *Server) Serve(network, addr string) error { } func (srv *Server) serve(ln net.Listener) error { + select { + case <-srv.stopC: + return errors.New("unable to serve because server is allready stopped") + default: + srv.mu.Lock() + srv.listeners = append(srv.listeners, ln) + srv.mu.Unlock() + } + for { + select { + case <-srv.stopC: + return nil + default: + // do nothing + } + cn, err := ln.Accept() if err != nil { log.Printf("Accept failed: %+v", err) diff --git a/ldap/server_test.go b/ldap/server_test.go new file mode 100644 index 0000000..8b543ad --- /dev/null +++ b/ldap/server_test.go @@ -0,0 +1,62 @@ +package ldap + +import ( + "crypto/tls" + "errors" + "github.com/stretchr/testify/assert" + "net" + "testing" + "time" +) + +func TestServer_Shutdown(t *testing.T) { + // try to start with closed listener + server := &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{})} + assert.NoError(t, server.Shutdown()) + assert.Equal(t, errors.New("unable to serve because server is allready stopped"), server.Serve("", "")) + + // shutdown while serving + server = &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{})} + time.AfterFunc(2*time.Second, func() { server.Shutdown() }) + assert.NoError(t, server.Serve("", "")) +} + +// listenerFactoryMock - listener factory mock for unit testing +var _ listenerFactory = (*listenerFactoryMock)(nil) + +type listenerFactoryMock struct{} + +func (_ *listenerFactoryMock) newListener(network, address string) (net.Listener, error) { + return newListenerMock(), nil +} + +func (_ *listenerFactoryMock) newTLSListener(network, addr string, config *tls.Config) (net.Listener, error) { + return newListenerMock(), nil +} + +// listenerMock - net.Listener mock for unit testing +var _ net.Listener = (*listenerMock)(nil) + +type listenerMock struct { + stopC chan struct{} +} + +func newListenerMock() *listenerMock { + return &listenerMock{ + stopC: make(chan struct{}), + } +} + +func (l *listenerMock) Accept() (net.Conn, error) { + <-l.stopC + return nil, errors.New("listener is closed") +} + +func (l *listenerMock) Close() error { + close(l.stopC) + return nil +} + +func (l *listenerMock) Addr() net.Addr { + return nil +} From 7628ab2daa61d121032d6e29e466b6cc28142107 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 15:43:37 +0300 Subject: [PATCH 2/8] ldap_server_shutdown | wait until all active connections will closed in Server.Shutdown --- ldap/server.go | 40 +++++++++++++++++++++- ldap/server_test.go | 83 ++++++++++++++++++++++++++++++++++++++++++--- 2 files changed, 117 insertions(+), 6 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index 9709ede..279c28c 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -11,6 +11,7 @@ import ( "os" "strings" "sync" + "sync/atomic" ) func NewResponsePacket(msgID int) *Packet { @@ -81,12 +82,43 @@ type Server struct { RootDSE map[string][]string tlsConfig *tls.Config + wg *waitGroup mu sync.Mutex + once sync.Once listeners []net.Listener stopC chan struct{} factory listenerFactory } +// waitGroup - kind of a std WaitGroup +// we can't use std WaitGroup here to keep original Server API +type waitGroup struct { + cond *sync.Cond + counter int32 +} + +func newWaitGroup() *waitGroup { + return &waitGroup{cond: sync.NewCond(&sync.Mutex{})} +} + +func (w *waitGroup) add() { + atomic.AddInt32(&w.counter, 1) +} + +func (w *waitGroup) done() { + atomic.AddInt32(&w.counter, -1) + w.cond.Broadcast() +} + +func (w *waitGroup) wait() { + w.cond.Broadcast() + w.cond.L.Lock() + for atomic.LoadInt32(&w.counter) > 0 { + w.cond.Wait() + } + w.cond.L.Unlock() +} + type srvClient struct { cn net.Conn wr *bufio.Writer @@ -127,6 +159,7 @@ func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) { tlsConfig: tlsConfig, factory: &listenerFactoryImpl{}, stopC: make(chan struct{}), + wg: newWaitGroup(), }, nil } @@ -142,6 +175,8 @@ func (srv *Server) Shutdown() error { } srv.mu.Unlock() + srv.wg.wait() + return err } @@ -170,13 +205,16 @@ func (srv *Server) Serve(network, addr string) error { func (srv *Server) serve(ln net.Listener) error { select { case <-srv.stopC: - return errors.New("unable to serve because server is allready stopped") + return nil default: srv.mu.Lock() srv.listeners = append(srv.listeners, ln) srv.mu.Unlock() } + srv.wg.add() + defer srv.wg.done() + for { select { case <-srv.stopC: diff --git a/ldap/server_test.go b/ldap/server_test.go index 8b543ad..810b8cd 100644 --- a/ldap/server_test.go +++ b/ldap/server_test.go @@ -4,21 +4,94 @@ import ( "crypto/tls" "errors" "github.com/stretchr/testify/assert" + "golang.org/x/net/context" "net" + "sync/atomic" "testing" "time" ) +func TestWaitGroup_Wait(t *testing.T) { + var ( + w = newWaitGroup() + readyC = make(chan struct{}) + stuck bool + ) + + go func() { + w.wait() + close(readyC) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + + select { + case <-ctx.Done(): + stuck = true + case <-readyC: + // ok + } + + assert.False(t, stuck) +} + +func TestWaitGroup_AddDone(t *testing.T) { + var ( + w = newWaitGroup() + n = 1000 + readyC = make(chan struct{}) + stuck bool + ) + + for i := 0; i < n; i++ { + go func() { + w.add() + time.Sleep(1 * time.Millisecond) + w.done() + }() + } + + go func() { + w.wait() + close(readyC) + }() + + ctx, cancel := context.WithTimeout(context.Background(), time.Duration(2*n)*time.Millisecond) + defer cancel() + + select { + case <-ctx.Done(): + stuck = true + case <-readyC: + // ok + } + + assert.False(t, stuck) + assert.Equal(t, int32(0), atomic.LoadInt32(&w.counter)) +} + func TestServer_Shutdown(t *testing.T) { // try to start with closed listener - server := &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{})} + server := &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{}), wg: newWaitGroup()} assert.NoError(t, server.Shutdown()) - assert.Equal(t, errors.New("unable to serve because server is allready stopped"), server.Serve("", "")) + assert.NoError(t, server.Serve("", "")) // shutdown while serving - server = &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{})} - time.AfterFunc(2*time.Second, func() { server.Shutdown() }) - assert.NoError(t, server.Serve("", "")) + server = &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{}), wg: newWaitGroup()} + n := 1000 + for i := 0; i < n; i++ { + go func() { + assert.NoError(t, server.Serve("", "")) + }() + } + + time.AfterFunc(2*time.Second, func() { + assert.Equal(t, int32(n), atomic.LoadInt32(&server.wg.counter)) + server.Shutdown() + assert.Equal(t, int32(0), atomic.LoadInt32(&server.wg.counter)) + + }) } // listenerFactoryMock - listener factory mock for unit testing From 01e5e696d99d7b4bc147742b361d7f4c57b84bd0 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 17:45:58 +0300 Subject: [PATCH 3/8] ldap_server_shutdown | call wg.add/done in connection handler --- ldap/server.go | 9 ++++++--- ldap/server_test.go | 25 ++++++++++++++++++++----- 2 files changed, 26 insertions(+), 8 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index 279c28c..264857b 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -92,6 +92,7 @@ type Server struct { // waitGroup - kind of a std WaitGroup // we can't use std WaitGroup here to keep original Server API +// std WaitGroup doesn't allow to call Add before Wait type waitGroup struct { cond *sync.Cond counter int32 @@ -121,6 +122,7 @@ func (w *waitGroup) wait() { type srvClient struct { cn net.Conn + wg *waitGroup wr *bufio.Writer srv *Server ctx Context @@ -212,9 +214,6 @@ func (srv *Server) serve(ln net.Listener) error { srv.mu.Unlock() } - srv.wg.add() - defer srv.wg.done() - for { select { case <-srv.stopC: @@ -231,6 +230,7 @@ func (srv *Server) serve(ln net.Listener) error { go (&srvClient{ cn: cn, + wg: srv.wg, wr: bufio.NewWriter(cn), srv: srv, }).serve() @@ -238,6 +238,9 @@ func (srv *Server) serve(ln net.Listener) error { } func (cli *srvClient) serve() { + cli.wg.add() + defer cli.wg.done() + ctx, err := cli.srv.Backend.Connect(cli.cn.RemoteAddr()) if err != nil { cli.cn.Close() diff --git a/ldap/server_test.go b/ldap/server_test.go index 810b8cd..665f1ad 100644 --- a/ldap/server_test.go +++ b/ldap/server_test.go @@ -77,21 +77,36 @@ func TestServer_Shutdown(t *testing.T) { assert.NoError(t, server.Shutdown()) assert.NoError(t, server.Serve("", "")) + var ( + readyC = make(chan struct{}) + stuck bool + ) + // shutdown while serving server = &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{}), wg: newWaitGroup()} - n := 1000 + n := 10 for i := 0; i < n; i++ { go func() { assert.NoError(t, server.Serve("", "")) }() } - time.AfterFunc(2*time.Second, func() { - assert.Equal(t, int32(n), atomic.LoadInt32(&server.wg.counter)) + time.AfterFunc(1*time.Second, func() { server.Shutdown() - assert.Equal(t, int32(0), atomic.LoadInt32(&server.wg.counter)) - + close(readyC) }) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + select { + case <-ctx.Done(): + stuck = true + case <-readyC: + // ok + } + + assert.False(t, stuck) } // listenerFactoryMock - listener factory mock for unit testing From 0c2ba145d17555ba56a2a5f771a2bfb02c35b96e Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 17:52:01 +0300 Subject: [PATCH 4/8] ldap_server_shutdown | fix comment + error checking --- ldap/server.go | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index 264857b..34af40a 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -92,7 +92,8 @@ type Server struct { // waitGroup - kind of a std WaitGroup // we can't use std WaitGroup here to keep original Server API -// std WaitGroup doesn't allow to call Add before Wait +// std WaitGroup doesn't allow to call Wait before Add +// but technically we can call Shutdown before Serve type waitGroup struct { cond *sync.Cond counter int32 @@ -171,7 +172,7 @@ func (srv *Server) Shutdown() error { close(srv.stopC) srv.mu.Lock() for _, listener := range srv.listeners { - if e := listener.Close(); err != nil { + if e := listener.Close(); e != nil { err = e } } From 029e51a072dd5593ec1aa10f64fcc6bdad8e4ec9 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 18:43:31 +0300 Subject: [PATCH 5/8] ldap_server_shutdown | stopC channel replaced with is isStopped flag --- ldap/server.go | 34 +++++++++++++++++----------------- ldap/server_test.go | 19 ++++++++++++++++--- 2 files changed, 33 insertions(+), 20 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index 34af40a..8033b2d 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -86,7 +86,7 @@ type Server struct { mu sync.Mutex once sync.Once listeners []net.Listener - stopC chan struct{} + isStopped int32 factory listenerFactory } @@ -161,7 +161,6 @@ func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) { RootDSE: sf, tlsConfig: tlsConfig, factory: &listenerFactoryImpl{}, - stopC: make(chan struct{}), wg: newWaitGroup(), }, nil } @@ -169,8 +168,8 @@ func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) { func (srv *Server) Shutdown() error { var err error - close(srv.stopC) srv.mu.Lock() + atomic.StoreInt32(&srv.isStopped, 1) for _, listener := range srv.listeners { if e := listener.Close(); e != nil { err = e @@ -206,27 +205,21 @@ func (srv *Server) Serve(network, addr string) error { } func (srv *Server) serve(ln net.Listener) error { - select { - case <-srv.stopC: - return nil - default: - srv.mu.Lock() - srv.listeners = append(srv.listeners, ln) + srv.mu.Lock() + if atomic.LoadInt32(&srv.isStopped) > 0 { srv.mu.Unlock() + return nil } + srv.listeners = append(srv.listeners, ln) + srv.mu.Unlock() for { - select { - case <-srv.stopC: - return nil - default: - // do nothing - } - cn, err := ln.Accept() - if err != nil { + if err != nil && isTemporary(err) { log.Printf("Accept failed: %+v", err) continue + } else { + return err } go (&srvClient{ @@ -238,6 +231,13 @@ func (srv *Server) serve(ln net.Listener) error { } } +func isTemporary(err error) bool { + if ne, ok := err.(net.Error); ok { + return ne.Temporary() + } + return false +} + func (cli *srvClient) serve() { cli.wg.add() defer cli.wg.done() diff --git a/ldap/server_test.go b/ldap/server_test.go index 665f1ad..c85e267 100644 --- a/ldap/server_test.go +++ b/ldap/server_test.go @@ -73,7 +73,7 @@ func TestWaitGroup_AddDone(t *testing.T) { func TestServer_Shutdown(t *testing.T) { // try to start with closed listener - server := &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{}), wg: newWaitGroup()} + server := &Server{factory: &listenerFactoryMock{}, wg: newWaitGroup()} assert.NoError(t, server.Shutdown()) assert.NoError(t, server.Serve("", "")) @@ -83,7 +83,7 @@ func TestServer_Shutdown(t *testing.T) { ) // shutdown while serving - server = &Server{factory: &listenerFactoryMock{}, stopC: make(chan struct{}), wg: newWaitGroup()} + server = &Server{factory: &listenerFactoryMock{}, wg: newWaitGroup()} n := 10 for i := 0; i < n; i++ { go func() { @@ -137,7 +137,7 @@ func newListenerMock() *listenerMock { func (l *listenerMock) Accept() (net.Conn, error) { <-l.stopC - return nil, errors.New("listener is closed") + return nil, &netErrorMock{errors.New("listener is closed")} } func (l *listenerMock) Close() error { @@ -148,3 +148,16 @@ func (l *listenerMock) Close() error { func (l *listenerMock) Addr() net.Addr { return nil } + +// netErrorMock - net.Error impl for unit testing +type netErrorMock struct { + error +} + +func (_ *netErrorMock) Temporary() bool { + return true +} + +func (_ *netErrorMock) Timeout() bool { + return true +} From 7dda82daf767396358b5f45b7f1370ef96425e5c Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 18:49:18 +0300 Subject: [PATCH 6/8] ldap_server_shutdown | change isStopped flag type from int32 to bool --- ldap/server.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index 8033b2d..e3e8e8e 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -86,7 +86,7 @@ type Server struct { mu sync.Mutex once sync.Once listeners []net.Listener - isStopped int32 + isStopped bool factory listenerFactory } @@ -169,7 +169,7 @@ func (srv *Server) Shutdown() error { var err error srv.mu.Lock() - atomic.StoreInt32(&srv.isStopped, 1) + srv.isStopped = true for _, listener := range srv.listeners { if e := listener.Close(); e != nil { err = e @@ -206,7 +206,7 @@ func (srv *Server) Serve(network, addr string) error { func (srv *Server) serve(ln net.Listener) error { srv.mu.Lock() - if atomic.LoadInt32(&srv.isStopped) > 0 { + if srv.isStopped { srv.mu.Unlock() return nil } From ac2bd9cf546971bbb3d991b802d3e37b6ac8d6b0 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Fri, 10 Nov 2017 18:51:53 +0300 Subject: [PATCH 7/8] ldap_server_shutdown | change error handling in Server.serve --- ldap/server.go | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/ldap/server.go b/ldap/server.go index e3e8e8e..8f2d6af 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -215,10 +215,11 @@ func (srv *Server) serve(ln net.Listener) error { for { cn, err := ln.Accept() - if err != nil && isTemporary(err) { - log.Printf("Accept failed: %+v", err) - continue - } else { + if err != nil { + if isTemporary(err) { + log.Printf("Accept failed: %+v", err) + continue + } return err } From 2557f64419ced7f323729c03c93276f118c8aed9 Mon Sep 17 00:00:00 2001 From: Denis Shilkin Date: Mon, 13 Nov 2017 10:02:04 +0300 Subject: [PATCH 8/8] ldap_server_shutdown | delete unused sync.Once --- ldap/server.go | 1 - 1 file changed, 1 deletion(-) diff --git a/ldap/server.go b/ldap/server.go index 8f2d6af..4534f4d 100644 --- a/ldap/server.go +++ b/ldap/server.go @@ -84,7 +84,6 @@ type Server struct { tlsConfig *tls.Config wg *waitGroup mu sync.Mutex - once sync.Once listeners []net.Listener isStopped bool factory listenerFactory