Skip to content
102 changes: 98 additions & 4 deletions ldap/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ import (
"net"
"os"
"strings"
"sync"
"sync/atomic"
)

func NewResponsePacket(msgID int) *Packet {
Expand Down Expand Up @@ -80,15 +82,66 @@ type Server struct {
RootDSE map[string][]string

tlsConfig *tls.Config
wg *waitGroup
mu sync.Mutex
listeners []net.Listener
isStopped bool
factory listenerFactory
}

// 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 Wait before Add
// but technically we can call Shutdown before Serve
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
wg *waitGroup
wr *bufio.Writer
srv *Server
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))
Expand All @@ -106,48 +159,89 @@ func NewServer(be Backend, tlsConfig *tls.Config) (*Server, error) {
Backend: be,
RootDSE: sf,
tlsConfig: tlsConfig,
factory: &listenerFactoryImpl{},
wg: newWaitGroup(),
}, nil
}

func (srv *Server) Shutdown() error {
var err error

srv.mu.Lock()
srv.isStopped = true
for _, listener := range srv.listeners {
if e := listener.Close(); e != nil {
err = e
}
}
srv.mu.Unlock()

srv.wg.wait()

return err
}

func (srv *Server) ServeTLS(network, addr string, tlsConfig *tls.Config) error {
if tlsConfig == nil {
tlsConfig = srv.tlsConfig
}
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
}
return srv.serve(ln)
}

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
}
return srv.serve(ln)
}

func (srv *Server) serve(ln net.Listener) error {
srv.mu.Lock()
if srv.isStopped {
srv.mu.Unlock()
return nil
}
srv.listeners = append(srv.listeners, ln)
srv.mu.Unlock()

for {
cn, err := ln.Accept()
if err != nil {
log.Printf("Accept failed: %+v", err)
continue
if isTemporary(err) {
log.Printf("Accept failed: %+v", err)
continue
}
return err
}

go (&srvClient{
cn: cn,
wg: srv.wg,
wr: bufio.NewWriter(cn),
srv: srv,
}).serve()
}
}

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()

ctx, err := cli.srv.Backend.Connect(cli.cn.RemoteAddr())
if err != nil {
cli.cn.Close()
Expand Down
163 changes: 163 additions & 0 deletions ldap/server_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
package ldap

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{}, wg: newWaitGroup()}
assert.NoError(t, server.Shutdown())
assert.NoError(t, server.Serve("", ""))

var (
readyC = make(chan struct{})
stuck bool
)

// shutdown while serving
server = &Server{factory: &listenerFactoryMock{}, wg: newWaitGroup()}
n := 10
for i := 0; i < n; i++ {
go func() {
assert.NoError(t, server.Serve("", ""))
}()
}

time.AfterFunc(1*time.Second, func() {
server.Shutdown()
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
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, &netErrorMock{errors.New("listener is closed")}
}

func (l *listenerMock) Close() error {
close(l.stopC)
return nil
}

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
}