From f643487ce086beb616a9ac00a500eb26df6f94f2 Mon Sep 17 00:00:00 2001 From: Matt Jaffee Date: Wed, 24 Oct 2018 16:55:39 -0500 Subject: [PATCH] add initial UDP proxy code to support proxying/partitioning memberlist --- internal/udproxy/udproxy.go | 145 +++++++++++++++++++++++++++++++ internal/udproxy/udproxy_test.go | 63 ++++++++++++++ 2 files changed, 208 insertions(+) create mode 100644 internal/udproxy/udproxy.go create mode 100644 internal/udproxy/udproxy_test.go diff --git a/internal/udproxy/udproxy.go b/internal/udproxy/udproxy.go new file mode 100644 index 000000000..fa30b22a5 --- /dev/null +++ b/internal/udproxy/udproxy.go @@ -0,0 +1,145 @@ +package udproxy + +import ( + "bytes" + "io" + "net" + "sync" + "time" + + "github.com/pkg/errors" + "golang.org/x/sync/errgroup" +) + +type Proxy struct { + upstreamAddr *net.UDPAddr + conn *net.UDPConn + + drop bool + dropLock sync.Mutex + + quit chan struct{} + eg errgroup.Group + // map from client address to upstream connection. We must maintain a + // separate connection to upstream for each client connection so that we can + // differentiate data sent back from upstream. + upstreams map[*net.UDPAddr]*net.UDPConn + // TODO - need to track a per-connection timeout so that "upstreams" doesn't + // grow indefinitely. +} + +func New(listenIP string, listenPort int, upstreamIP string, upstreamPort int) (*Proxy, error) { + uc, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP(listenIP), Port: listenPort}) + if err != nil { + return nil, errors.Wrap(err, "listening") + } + p := &Proxy{ + conn: uc, + upstreamAddr: &net.UDPAddr{IP: net.ParseIP(upstreamIP), Port: upstreamPort}, + quit: make(chan struct{}), + upstreams: make(map[*net.UDPAddr]*net.UDPConn), + } + if p.upstreamAddr.IP == nil { + return nil, errors.Errorf("unable to parse upstream ip '%s'", upstreamIP) + } + p.eg.Go(p.run) + return p, nil +} + +func (p *Proxy) Drop() { + p.dropLock.Lock() + p.drop = true + p.dropLock.Unlock() +} + +func (p *Proxy) Undrop() { + p.dropLock.Lock() + p.drop = false + p.dropLock.Unlock() +} + +func (p *Proxy) dropping() bool { + p.dropLock.Lock() + d := p.drop + p.dropLock.Unlock() + return d +} + +func (p *Proxy) run() error { + buf := make([]byte, 65507) + for { + select { + case <-p.quit: + return nil + default: + } + err := p.conn.SetReadDeadline(time.Now().Add(time.Millisecond * 10)) + if err != nil { + return errors.Wrap(err, "setting read deadline (run)") + } + n, addr, err := p.conn.ReadFromUDP(buf) + if err, ok := err.(net.Error); ok && err.Timeout() { + continue + } else if err != nil { + return errors.Wrap(err, "reading from udp conn") + } + upConn := p.upstreams[addr] + if upConn == nil { + p.upstreams[addr], err = net.DialUDP("udp", &net.UDPAddr{}, p.upstreamAddr) + if err != nil { + return errors.Wrap(err, "creating new connection to upstream") + } + p.eg.Go(func() error { + return p.proxyBack(addr, p.upstreams[addr]) + }) + upConn = p.upstreams[addr] + } + if !p.dropping() { + _, err = io.Copy(upConn, bytes.NewBuffer(buf[:n])) + if err != nil { + return errors.Wrap(err, "writing to upstream conn") + } + } + } +} + +func (p *Proxy) proxyBack(to *net.UDPAddr, from *net.UDPConn) error { + buf := make([]byte, 65507) + for { + select { + case <-p.quit: + return nil + default: + } + err := from.SetReadDeadline(time.Now().Add(time.Millisecond * 10)) + if err != nil { + return errors.Wrap(err, "setting read deadline (proxyBack)") + } + n, _, err := from.ReadFromUDP(buf) + if err, ok := err.(net.Error); ok && err.Timeout() { + continue + } else if err != nil { + return errors.Wrap(err, "reading from upstream") + } + if !p.dropping() { + _, err = io.Copy(addrWriter{c: p.conn, a: to}, bytes.NewBuffer(buf[:n])) + if err != nil { + return errors.Wrap(err, "writing back to client") + } + } + } +} + +type addrWriter struct { + c *net.UDPConn + a *net.UDPAddr +} + +func (a addrWriter) Write(b []byte) (n int, err error) { + return a.c.WriteTo(b, a.a) +} + +func (p *Proxy) Close() error { + close(p.quit) + return p.eg.Wait() +} diff --git a/internal/udproxy/udproxy_test.go b/internal/udproxy/udproxy_test.go new file mode 100644 index 000000000..edb915d1f --- /dev/null +++ b/internal/udproxy/udproxy_test.go @@ -0,0 +1,63 @@ +package udproxy_test + +import ( + "net" + "testing" + + "github.com/pilosa/pilosa/internal/udproxy" + "github.com/pkg/errors" + "golang.org/x/sync/errgroup" +) + +func TestUDProxy(t *testing.T) { + p, err := udproxy.New("127.0.0.1", 12345, "127.0.0.1", 12346) + if err != nil { + t.Fatalf("creating proxy: %v", err) + } + + uc, err := net.ListenUDP("udp", &net.UDPAddr{Port: 12346}) + if err != nil { + t.Fatalf("listening udp upstream: %v", err) + } + + resp := make([]byte, 8) + eg := errgroup.Group{} + eg.Go(func() error { + conn, err := net.DialUDP("udp", nil, &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 12345}) + if err != nil { + t.Fatalf("connecting to proxy: %v", err) + } + _, err = conn.Write([]byte("hello!")) + if err != nil { + return errors.Wrap(err, "writing to proxy") + } + _, err = conn.Read(resp) + if err != nil { + return errors.Wrap(err, "reading from proxy") + } + return nil + }) + + req := make([]byte, 10) + _, addr, err := uc.ReadFrom(req) + if err != nil { + t.Fatalf("upstream reading from proxy: %v", err) + } + if string(req[:6]) != "hello!" { + t.Fatalf("got unexpected request %s", req) + } + _, err = uc.WriteTo([]byte("goodbye"), addr) + if err != nil { + t.Fatalf("writing response: %v", err) + } + eg.Wait() + if string(resp[:7]) != "goodbye" { + t.Fatalf("got unexpected response '%v", resp) + } + err = p.Close() + if err != nil { + t.Fatalf("err closing proxy: '%v'", err) + } +} + +// TODO test dropping