From 35f704e2e488646de4aa209c892a4a87183daf7a Mon Sep 17 00:00:00 2001 From: David Zhao Date: Sun, 18 Apr 2021 22:27:37 -0700 Subject: [PATCH] Fix UDPMux read/write race condition Resolves #351 --- agent_udpmux_test.go | 10 ++++- udp_mux.go | 103 +++++++++++++++---------------------------- udp_mux_test.go | 55 ++++++++++++++++++----- udp_muxed_conn.go | 52 ++++++++++++++++------ 4 files changed, 126 insertions(+), 94 deletions(-) diff --git a/agent_udpmux_test.go b/agent_udpmux_test.go index e7b52e1..b561d06 100644 --- a/agent_udpmux_test.go +++ b/agent_udpmux_test.go @@ -3,6 +3,7 @@ package ice import ( + "net" "testing" "time" @@ -24,7 +25,14 @@ func TestMuxAgent(t *testing.T) { Logger: loggerFactory.NewLogger("ice"), }) muxPort := 7686 - require.NoError(t, udpMux.Start(muxPort)) + c, err := net.ListenUDP(udp, &net.UDPAddr{ + Port: muxPort, + }) + require.NoError(t, err) + require.NoError(t, udpMux.Start(c)) + defer func() { + _ = udpMux.Close() + }() muxedA, err := NewAgent(&AgentConfig{ UDPMux: udpMux, diff --git a/udp_mux.go b/udp_mux.go index 42c2240..f8688b0 100644 --- a/udp_mux.go +++ b/udp_mux.go @@ -8,7 +8,6 @@ import ( "os" "strings" "sync" - "time" "github.com/pion/logging" ) @@ -18,34 +17,30 @@ type UDPMux interface { io.Closer GetConn(ufrag, network string) (net.PacketConn, error) RemoveConnByUfrag(ufrag string) - Start(port int) error + Start(conn net.PacketConn) error } // UDPMuxDefault is an implementation of the interface type UDPMuxDefault struct { - params UDPMuxParams - listenAddr *net.UDPAddr - udpConn *net.UDPConn + params UDPMuxParams + udpConn net.PacketConn - mappingChan chan connMap - closedChan chan struct{} - closeOnce sync.Once + closedChan chan struct{} + closeOnce sync.Once // conns is a map of all udpMuxedConn indexed by ufrag|network|candidateType conns map[string]*udpMuxedConn + // map of udpAddr -> udpMuxedConn + addressMap sync.Map + // 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 -} +const maxAddrSize = 512 // UDPMuxParams are parameters for UDPMux. type UDPMuxParams struct { @@ -55,31 +50,26 @@ type UDPMuxParams struct { // NewUDPMuxDefault creates an implementation of UDPMux func NewUDPMuxDefault(params UDPMuxParams) *UDPMuxDefault { return &UDPMuxDefault{ - params: params, - conns: make(map[string]*udpMuxedConn), - mappingChan: make(chan connMap, 10), - closedChan: make(chan struct{}, 1), + params: params, + conns: make(map[string]*udpMuxedConn), + closedChan: make(chan struct{}, 1), pool: &sync.Pool{ New: func() interface{} { - return newBufferHolder(maxAddrSize) + // big enough buffer to fit both packet and address + return newBufferHolder(receiveMTU + maxAddrSize) }, }, } } // Start starts the mux. Before the UDPMux is usable, it must be started -func (m *UDPMuxDefault) Start(port int) error { +// the mux will read/write data on the underlying net.PacketConn. It must +// conn.ReadFrom *MUST* return a *net.UDPAddr +func (m *UDPMuxDefault) Start(conn net.PacketConn) error { if m.udpConn != nil { return ErrMultipleStart } - m.listenAddr = &net.UDPAddr{ - Port: port, - } - uc, err := net.ListenUDP(udp, m.listenAddr) - if err != nil { - return err - } - m.udpConn = uc + m.udpConn = conn go m.connWorker() return nil @@ -87,7 +77,7 @@ func (m *UDPMuxDefault) Start(port int) error { // LocalAddr returns the listening address of this UDPMuxDefault func (m *UDPMuxDefault) LocalAddr() net.Addr { - return m.listenAddr + return m.udpConn.LocalAddr() } // GetConn returns a PacketConn given the connection's ufrag and network @@ -139,10 +129,7 @@ func (m *UDPMuxDefault) RemoveConnByUfrag(ufrag string) { for _, c := range removedConns { addresses := c.getAddresses() for _, addr := range addresses { - m.mappingChan <- connMap{ - address: addr, - conn: nil, - } + m.addressMap.Delete(addr) } } } @@ -189,10 +176,7 @@ func (m *UDPMuxDefault) removeConn(key string) { addresses := c.getAddresses() for _, addr := range addresses { - m.mappingChan <- connMap{ - address: addr, - conn: nil, - } + m.addressMap.Delete(addr) } } @@ -204,10 +188,11 @@ func (m *UDPMuxDefault) registerConnForAddress(conn *udpMuxedConn, addr string) if m.IsClosed() { return } - m.mappingChan <- connMap{ - address: addr, - conn: conn, + existing, ok := m.addressMap.Load(addr) + if ok { + existing.(*udpMuxedConn).removeAddress(addr) } + m.addressMap.Store(addr, conn) } func (m *UDPMuxDefault) createMuxedConn(key string) *udpMuxedConn { @@ -222,9 +207,6 @@ func (m *UDPMuxDefault) createMuxedConn(key string) *udpMuxedConn { } func (m *UDPMuxDefault) connWorker() { - // map of remote addresses -> udpMuxedConn - // used to look up incoming packets - remoteMap := make(map[string]*udpMuxedConn) logger := m.params.Logger defer func() { @@ -232,10 +214,7 @@ func (m *UDPMuxDefault) connWorker() { }() buf := make([]byte, receiveMTU) for { - _ = m.udpConn.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) - n, addr, err := m.udpConn.ReadFromUDP(buf) - // process any mapping changes, this is done as early as possible - m.applyMappingChanges(remoteMap) + n, addr, err := m.udpConn.ReadFrom(buf) if err != nil { if errors.Is(err, os.ErrDeadlineExceeded) { continue @@ -246,37 +225,25 @@ func (m *UDPMuxDefault) connWorker() { } // look up forward destination - addrStr := addr.String() - c := remoteMap[addrStr] - - if c == nil { + v, ok := m.addressMap.Load(addr.String()) + if !ok { // ignore packets that we don't know where to route to continue } - err = c.writePacket(buf[:n], addr) + udpAddr, ok := addr.(*net.UDPAddr) + if !ok { + logger.Errorf("underlying PacketConn did not return a UDPAddr") + return + } + c := v.(*udpMuxedConn) + err = c.writePacket(buf[:n], udpAddr) if err != nil { logger.Errorf("could not write packet: %v", err) } } } -func (m *UDPMuxDefault) applyMappingChanges(remoteMap map[string]*udpMuxedConn) { - for { - select { - case cm := <-m.mappingChan: - // deregister previous addresses - existingConn := remoteMap[cm.address] - if existingConn != nil { - existingConn.removeAddress(cm.address) - } - remoteMap[cm.address] = cm.conn - default: - return - } - } -} - type bufferHolder struct { buffer []byte } diff --git a/udp_mux_test.go b/udp_mux_test.go index 0fb982c..192b85e 100644 --- a/udp_mux_test.go +++ b/udp_mux_test.go @@ -29,7 +29,10 @@ func TestUDPMux(t *testing.T) { udpMux := NewUDPMuxDefault(UDPMuxParams{ Logger: loggerFactory.NewLogger("ice"), }) - err := udpMux.Start(7686) + + conn, err := net.ListenUDP(udp, &net.UDPAddr{}) + require.NoError(t, err) + err = udpMux.Start(conn) require.NoError(t, err) defer func() { @@ -37,28 +40,32 @@ func TestUDPMux(t *testing.T) { }() require.NotNil(t, udpMux.LocalAddr(), "tcpMux.LocalAddr() is nil") - require.Equal(t, ":7686", udpMux.LocalAddr().String()) wg := sync.WaitGroup{} wg.Add(1) go func() { defer wg.Done() - testMuxConnection(t, udpMux, "ufrag1") + testMuxConnection(t, udpMux, "ufrag1", udp) }() wg.Add(1) go func() { defer wg.Done() - testMuxConnection(t, udpMux, "ufrag2") + testMuxConnection(t, udpMux, "ufrag2", "udp4") }() - testMuxConnection(t, udpMux, "ufrag3") + // skip ipv6 test on i386 + const ptrSize = 32 << (^uintptr(0) >> 63) + if ptrSize != 32 { + testMuxConnection(t, udpMux, "ufrag3", "udp6") + } + wg.Wait() require.NoError(t, udpMux.Close()) // can't create more connections - _, err = udpMux.GetConn("failufrag", "udp") + _, err = udpMux.GetConn("failufrag", udp) require.Error(t, err) } @@ -102,14 +109,16 @@ func TestAddressEncoding(t *testing.T) { } } -func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { - pktConn, err := udpMux.GetConn(ufrag, udp) +func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string, network string) { + pktConn, err := udpMux.GetConn(ufrag, network) require.NoError(t, err, "error retrieving muxed connection for ufrag") defer func() { _ = pktConn.Close() }() - remoteConn, err := net.DialUDP(udp, nil, udpMux.LocalAddr().(*net.UDPAddr)) + remoteConn, err := net.DialUDP(network, nil, &net.UDPAddr{ + Port: udpMux.LocalAddr().(*net.UDPAddr).Port, + }) require.NoError(t, err, "error dialing test udp connection") // initial messages are dropped @@ -135,6 +144,7 @@ func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { // start writing packets through mux targetSize := 1 * 1024 * 1024 readDone := make(chan struct{}, 1) + remoteReadDone := make(chan struct{}, 1) // read packets from the muxed side go func() { @@ -145,11 +155,33 @@ func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { readBuf := make([]byte, receiveMTU) nextSeq := uint32(0) for read := 0; read < targetSize; { - n, _, _ := pktConn.ReadFrom(readBuf) + n, _, err := pktConn.ReadFrom(readBuf) require.NoError(t, err) require.Equal(t, receiveMTU, n) - verifyPacket(t, readBuf, nextSeq) + verifyPacket(t, readBuf[:n], nextSeq) + + // write it back to sender + _, err = pktConn.WriteTo(readBuf[:n], remoteConn.LocalAddr()) + require.NoError(t, err) + + read += n + nextSeq++ + } + }() + + go func() { + defer func() { + close(remoteReadDone) + }() + readBuf := make([]byte, receiveMTU) + nextSeq := uint32(0) + for read := 0; read < targetSize; { + n, _, err := remoteConn.ReadFrom(readBuf) + require.NoError(t, err) + require.Equal(t, receiveMTU, n) + + verifyPacket(t, readBuf[:n], nextSeq) read += n nextSeq++ @@ -178,6 +210,7 @@ func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) { } <-readDone + <-remoteReadDone } func verifyPacket(t *testing.T, b []byte, nextSeq uint32) { diff --git a/udp_muxed_conn.go b/udp_muxed_conn.go index 66a8f18..2c71238 100644 --- a/udp_muxed_conn.go +++ b/udp_muxed_conn.go @@ -47,22 +47,30 @@ func (c *udpMuxedConn) ReadFrom(b []byte) (n int, raddr net.Addr, err error) { defer c.params.AddrPool.Put(buf) // read address - addrN, err := c.buffer.Read(buf.buffer) + total, err := c.buffer.Read(buf.buffer) if err != nil { return 0, nil, err } - if raddr, err = decodeUDPAddr(buf.buffer[:addrN]); err != nil { + dataLen := int(binary.LittleEndian.Uint16(buf.buffer[:2])) + if dataLen > total || dataLen > len(b) { + return 0, nil, io.ErrShortBuffer + } + + // read data and then address + offset := 2 + copy(b, buf.buffer[offset:offset+dataLen]) + offset += dataLen + + // read address len & decode address + addrLen := int(binary.LittleEndian.Uint16(buf.buffer[offset : offset+2])) + offset += 2 + + if raddr, err = decodeUDPAddr(buf.buffer[offset : offset+addrLen]); err != nil { return 0, nil, err } - // read data - n, err = c.buffer.Read(b) - if err != nil { - return 0, nil, err - } - - return n, raddr, err + return dataLen, raddr, nil } func (c *udpMuxedConn) WriteTo(buf []byte, raddr net.Addr) (n int, err error) { @@ -164,14 +172,30 @@ 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) + + // format of buffer | data len | data bytes | addr len | addr bytes | + if len(buf.buffer) < len(data)+maxAddrSize { + return io.ErrShortBuffer + } + // data len + binary.LittleEndian.PutUint16(buf.buffer, uint16(len(data))) + offset := 2 + + // data + copy(buf.buffer[offset:], data) + offset += len(data) + + // write address first, leaving room for its length + n, err := encodeUDPAddr(addr, buf.buffer[offset+2:]) if err != nil { return nil } - if _, err := c.buffer.Write(buf.buffer[:n]); err != nil { - return err - } - if _, err := c.buffer.Write(data); err != nil { + total := offset + n + 2 + + // address len + binary.LittleEndian.PutUint16(buf.buffer[offset:], uint16(n)) + + if _, err := c.buffer.Write(buf.buffer[:total]); err != nil { return err } return nil