Fix UDPMux read/write race condition

Resolves #351
This commit is contained in:
David Zhao
2021-04-19 18:36:38 -07:00
committed by Sean DuBois
parent 166ba31563
commit 35f704e2e4
4 changed files with 126 additions and 94 deletions
+9 -1
View File
@@ -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,
+35 -68
View File
@@ -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
}
+44 -11
View File
@@ -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) {
+38 -14
View File
@@ -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