diff --git a/agent_udpmux_test.go b/agent_udpmux_test.go index b403982..e7b52e1 100644 --- a/agent_udpmux_test.go +++ b/agent_udpmux_test.go @@ -21,8 +21,7 @@ func TestMuxAgent(t *testing.T) { loggerFactory := logging.NewDefaultLoggerFactory() udpMux := NewUDPMuxDefault(UDPMuxParams{ - Logger: loggerFactory.NewLogger("ice"), - ReadBufferSize: 20, + Logger: loggerFactory.NewLogger("ice"), }) muxPort := 7686 require.NoError(t, udpMux.Start(muxPort)) diff --git a/udp_mux.go b/udp_mux.go index 8b1a99a..42c2240 100644 --- a/udp_mux.go +++ b/udp_mux.go @@ -34,12 +34,14 @@ type UDPMuxDefault struct { // conns is a map of all udpMuxedConn indexed by ufrag|network|candidateType conns map[string]*udpMuxedConn - // buffer pool to recycle buffers for incoming packets + // buffer pool to recycle buffers for net.UDPAddr encodes/decodes pool *sync.Pool mu sync.Mutex } +const maxAddrSize = 256 + type connMap struct { address string conn *udpMuxedConn @@ -47,8 +49,7 @@ type connMap struct { // UDPMuxParams are parameters for UDPMux. type UDPMuxParams struct { - Logger logging.LeveledLogger - ReadBufferSize int + Logger logging.LeveledLogger } // NewUDPMuxDefault creates an implementation of UDPMux @@ -60,7 +61,7 @@ func NewUDPMuxDefault(params UDPMuxParams) *UDPMuxDefault { closedChan: make(chan struct{}, 1), pool: &sync.Pool{ New: func() interface{} { - return newBufferHolder(receiveMTU) + return newBufferHolder(maxAddrSize) }, }, } @@ -109,7 +110,7 @@ func (m *UDPMuxDefault) GetConn(ufrag, network string) (net.PacketConn, error) { return c, nil } - c := m.createMuxedConn() + c := m.createMuxedConn(key) go func() { <-c.CloseChannel() m.removeConn(key) @@ -199,10 +200,6 @@ func (m *UDPMuxDefault) writeTo(buf []byte, raddr net.Addr) (n int, err error) { return m.udpConn.WriteTo(buf, raddr) } -func (m *UDPMuxDefault) doneWithBuffer(buf *bufferHolder) { - m.pool.Put(buf) -} - func (m *UDPMuxDefault) registerConnForAddress(conn *udpMuxedConn, addr string) { if m.IsClosed() { return @@ -213,12 +210,13 @@ func (m *UDPMuxDefault) registerConnForAddress(conn *udpMuxedConn, addr string) } } -func (m *UDPMuxDefault) createMuxedConn() *udpMuxedConn { +func (m *UDPMuxDefault) createMuxedConn(key string) *udpMuxedConn { c := newUDPMuxedConn(&udpMuxedConnParams{ - Mux: m, - ReadBuffer: m.params.ReadBufferSize, - LocalAddr: m.LocalAddr(), - Logger: m.params.Logger, + Mux: m, + Key: key, + AddrPool: m.pool, + LocalAddr: m.LocalAddr(), + Logger: m.params.Logger, }) return c } @@ -232,15 +230,14 @@ func (m *UDPMuxDefault) connWorker() { defer func() { _ = m.Close() }() + buf := make([]byte, receiveMTU) for { - buffer := m.pool.Get().(*bufferHolder) _ = m.udpConn.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) - n, addr, err := m.udpConn.ReadFrom(buffer.buffer) - // process any mapping changes, this is done as early as possible to prevent channel clogging up + n, addr, err := m.udpConn.ReadFromUDP(buf) + // process any mapping changes, this is done as early as possible m.applyMappingChanges(remoteMap) if err != nil { if errors.Is(err, os.ErrDeadlineExceeded) { - m.doneWithBuffer(buffer) continue } else if err != io.EOF { logger.Errorf("could not read udp packet: %v", err) @@ -253,16 +250,11 @@ func (m *UDPMuxDefault) connWorker() { c := remoteMap[addrStr] if c == nil { - m.doneWithBuffer(buffer) // ignore packets that we don't know where to route to continue } - err = c.writePacket(muxedPacket{ - Buffer: buffer, - Size: n, - RAddr: addr, - }) + err = c.writePacket(buf[:n], addr) if err != nil { logger.Errorf("could not write packet: %v", err) } diff --git a/udp_mux_test.go b/udp_mux_test.go index bcce4d2..0fb982c 100644 --- a/udp_mux_test.go +++ b/udp_mux_test.go @@ -2,7 +2,11 @@ package ice +//nolint:gosec import ( + "crypto/rand" + "crypto/sha1" + "encoding/binary" "net" "sync" "testing" @@ -18,10 +22,12 @@ func TestUDPMux(t *testing.T) { report := test.CheckRoutines(t) defer report() + lim := test.TimeOut(time.Second * 30) + defer lim.Stop() + loggerFactory := logging.NewDefaultLoggerFactory() udpMux := NewUDPMuxDefault(UDPMuxParams{ - Logger: loggerFactory.NewLogger("ice"), - ReadBufferSize: 20, + Logger: loggerFactory.NewLogger("ice"), }) err := udpMux.Start(7686) require.NoError(t, err) @@ -56,6 +62,46 @@ func TestUDPMux(t *testing.T) { require.Error(t, err) } +func TestAddressEncoding(t *testing.T) { + cases := []struct { + name string + addr net.UDPAddr + }{ + { + name: "empty address", + }, + { + name: "ipv4", + addr: net.UDPAddr{ + IP: net.IPv4(244, 120, 0, 5), + Port: 6000, + Zone: "", + }, + }, + { + name: "ipv6", + addr: net.UDPAddr{ + IP: net.IPv6loopback, + Port: 2500, + Zone: "zone", + }, + }, + } + + for _, c := range cases { + addr := c.addr + t.Run(c.name, func(t *testing.T) { + buf := make([]byte, maxAddrSize) + n, err := encodeUDPAddr(&addr, buf) + require.NoError(t, err) + + parsedAddr, err := decodeUDPAddr(buf[:n]) + require.NoError(t, err) + require.EqualValues(t, &addr, parsedAddr) + }) + } +} + func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { pktConn, err := udpMux.GetConn(ufrag, udp) require.NoError(t, err, "error retrieving muxed connection for ufrag") @@ -86,22 +132,57 @@ func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { require.NoError(t, err) require.Equal(t, msg.Raw, buf[:n]) - // write a bunch of packets from remote to ensure proper receipt - dataToSend := [][]byte{ - []byte("hello world"), - []byte("test text"), - msg.Raw, - } + // start writing packets through mux + targetSize := 1 * 1024 * 1024 + readDone := make(chan struct{}, 1) - buffer := make([]byte, receiveMTU) - for _, data := range dataToSend { - _, err := remoteConn.Write(data) + // read packets from the muxed side + go func() { + defer func() { + t.Logf("closing read chan for: %s", ufrag) + close(readDone) + }() + readBuf := make([]byte, receiveMTU) + nextSeq := uint32(0) + for read := 0; read < targetSize; { + n, _, _ := pktConn.ReadFrom(readBuf) + require.NoError(t, err) + require.Equal(t, receiveMTU, n) + + verifyPacket(t, readBuf, nextSeq) + + read += n + nextSeq++ + } + }() + + sequence := 0 + for written := 0; written < targetSize; { + buf := make([]byte, receiveMTU) + // byte0-4: sequence + // bytes4-24: sha1 checksum + // bytes24-mtu: random data + _, err := rand.Read(buf[24:]) + require.NoError(t, err) + h := sha1.Sum(buf[24:]) //nolint:gosec + copy(buf[4:24], h[:]) + binary.LittleEndian.PutUint32(buf[0:4], uint32(sequence)) + + _, err = remoteConn.Write(buf) require.NoError(t, err) - n, _, err := pktConn.ReadFrom(buffer) - require.NoError(t, err) - require.Equal(t, data, buffer[:n]) + written += len(buf) + sequence++ - time.Sleep(10 * time.Millisecond) + time.Sleep(time.Millisecond) } + + <-readDone +} + +func verifyPacket(t *testing.T, b []byte, nextSeq uint32) { + readSeq := binary.LittleEndian.Uint32(b[0:4]) + require.Equal(t, nextSeq, readSeq) + h := sha1.Sum(b[24:]) //nolint:gosec + require.Equal(t, h[:], b[4:24]) } diff --git a/udp_muxed_conn.go b/udp_muxed_conn.go index fe439dc..66a8f18 100644 --- a/udp_muxed_conn.go +++ b/udp_muxed_conn.go @@ -1,25 +1,22 @@ package ice import ( + "encoding/binary" "io" "net" "sync" "time" "github.com/pion/logging" + "github.com/pion/transport/packetio" ) type udpMuxedConnParams struct { - Mux *UDPMuxDefault - ReadBuffer int - LocalAddr net.Addr - Logger logging.LeveledLogger -} - -type muxedPacket struct { - Buffer *bufferHolder - RAddr net.Addr - Size int + Mux *UDPMuxDefault + AddrPool *sync.Pool + Key string + LocalAddr net.Addr + Logger logging.LeveledLogger } // udpMuxedConn represents a logical packet conn for a single remote as identified by ufrag @@ -29,7 +26,7 @@ type udpMuxedConn struct { addresses []string // channel holding incoming packets - recvChan chan muxedPacket + buffer *packetio.Buffer closedChan chan struct{} closeOnce sync.Once mu sync.Mutex @@ -38,7 +35,7 @@ type udpMuxedConn struct { func newUDPMuxedConn(params *udpMuxedConnParams) *udpMuxedConn { p := &udpMuxedConn{ params: params, - recvChan: make(chan muxedPacket, params.ReadBuffer), + buffer: packetio.NewBuffer(), closedChan: make(chan struct{}), } @@ -46,19 +43,26 @@ func newUDPMuxedConn(params *udpMuxedConnParams) *udpMuxedConn { } func (c *udpMuxedConn) ReadFrom(b []byte) (n int, raddr net.Addr, err error) { - pkt, ok := <-c.recvChan + buf := c.params.AddrPool.Get().(*bufferHolder) + defer c.params.AddrPool.Put(buf) - if !ok { - return 0, nil, io.ErrClosedPipe + // read address + addrN, err := c.buffer.Read(buf.buffer) + if err != nil { + return 0, nil, err } - if cap(b) < pkt.Size { - return 0, pkt.RAddr, io.ErrShortBuffer + if raddr, err = decodeUDPAddr(buf.buffer[:addrN]); err != nil { + return 0, nil, err } - copy(b, pkt.Buffer.buffer[:pkt.Size]) - c.params.Mux.doneWithBuffer(pkt.Buffer) - return pkt.Size, pkt.RAddr, err + // read data + n, err = c.buffer.Read(b) + if err != nil { + return 0, nil, err + } + + return n, raddr, err } func (c *udpMuxedConn) WriteTo(buf []byte, raddr net.Addr) (n int, err error) { @@ -95,14 +99,15 @@ func (c *udpMuxedConn) CloseChannel() <-chan struct{} { } func (c *udpMuxedConn) Close() error { + var err error c.closeOnce.Do(func() { + err = c.buffer.Close() close(c.closedChan) - close(c.recvChan) }) c.mu.Lock() defer c.mu.Unlock() c.addresses = nil - return nil + return err } func (c *udpMuxedConn) isClosed() bool { @@ -132,6 +137,8 @@ func (c *udpMuxedConn) addAddress(addr string) { } func (c *udpMuxedConn) removeAddress(addr string) { + c.mu.Lock() + defer c.mu.Unlock() newAddresses := make([]string, 0, len(c.addresses)) for _, a := range c.addresses { if a != addr { @@ -139,9 +146,7 @@ func (c *udpMuxedConn) removeAddress(addr string) { } } - c.mu.Lock() c.addresses = newAddresses - c.mu.Unlock() } func (c *udpMuxedConn) containsAddress(addr string) bool { @@ -155,11 +160,62 @@ func (c *udpMuxedConn) containsAddress(addr string) bool { return false } -func (c *udpMuxedConn) writePacket(pkt muxedPacket) error { - select { - case c.recvChan <- pkt: +func (c *udpMuxedConn) writePacket(data []byte, addr *net.UDPAddr) error { + // write two packets, address and data + buf := c.params.AddrPool.Get().(*bufferHolder) + defer c.params.AddrPool.Put(buf) + n, err := encodeUDPAddr(addr, buf.buffer) + if err != nil { return nil - case <-c.closedChan: - return io.ErrClosedPipe } + if _, err := c.buffer.Write(buf.buffer[:n]); err != nil { + return err + } + if _, err := c.buffer.Write(data); err != nil { + return err + } + return nil +} + +func encodeUDPAddr(addr *net.UDPAddr, buf []byte) (int, error) { + ipdata, err := addr.IP.MarshalText() + if err != nil { + return 0, err + } + total := 2 + len(ipdata) + 2 + len(addr.Zone) + if total > len(buf) { + return 0, io.ErrShortBuffer + } + + binary.LittleEndian.PutUint16(buf, uint16(len(ipdata))) + offset := 2 + n := copy(buf[offset:], ipdata) + offset += n + binary.LittleEndian.PutUint16(buf[offset:], uint16(addr.Port)) + offset += 2 + copy(buf[offset:], addr.Zone) + return total, nil +} + +func decodeUDPAddr(buf []byte) (*net.UDPAddr, error) { + addr := net.UDPAddr{} + + offset := 0 + ipLen := int(binary.LittleEndian.Uint16(buf[:2])) + offset += 2 + // basic bounds checking + if ipLen+offset > len(buf) { + return nil, io.ErrShortBuffer + } + if err := addr.IP.UnmarshalText(buf[offset : offset+ipLen]); err != nil { + return nil, err + } + offset += ipLen + addr.Port = int(binary.LittleEndian.Uint16(buf[offset : offset+2])) + offset += 2 + zone := make([]byte, len(buf[offset:])) + copy(zone, buf[offset:]) + addr.Zone = string(zone) + + return &addr, nil }