mirror of
https://github.com/netbirdio/ice-old.git
synced 2026-05-22 17:08:24 -07:00
Improved performance of UDPMux
The previous implementation of UDPMux was dropping packets due to channel usage in the critical write path.
This commit is contained in:
@@ -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))
|
||||
|
||||
+16
-24
@@ -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)
|
||||
}
|
||||
|
||||
+96
-15
@@ -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])
|
||||
}
|
||||
|
||||
+85
-29
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user