Improve UDPMux performance

improved buffer handling and prevents channel clogging
This commit is contained in:
David Zhao
2021-04-14 14:19:07 -07:00
parent 86d69d6ce5
commit f7b11daf96
5 changed files with 97 additions and 45 deletions
+5 -1
View File
@@ -4,6 +4,7 @@ package ice
import (
"testing"
"time"
"github.com/pion/logging"
"github.com/pion/transport/test"
@@ -14,8 +15,11 @@ import (
func TestMuxAgent(t *testing.T) {
report := test.CheckRoutines(t)
defer report()
loggerFactory := logging.NewDefaultLoggerFactory()
lim := test.TimeOut(time.Second * 30)
defer lim.Stop()
loggerFactory := logging.NewDefaultLoggerFactory()
udpMux := NewUDPMuxDefault(UDPMuxParams{
Logger: loggerFactory.NewLogger("ice"),
ReadBufferSize: 20,
+3 -1
View File
@@ -178,7 +178,7 @@ func (a *Agent) gatherCandidatesLocal(ctx context.Context, networkTypes []Networ
// accessible from the current interface.
case udp:
if a.udpMux != nil {
conn, err = a.udpMux.GetConnByUfrag(a.localUfrag)
conn, err = a.udpMux.GetConn(a.localUfrag, network)
if err != nil {
a.log.Warnf("could not get udp muxed connection: %v\n", err)
continue
@@ -237,6 +237,7 @@ func (a *Agent) gatherCandidatesSrflxMapped(ctx context.Context, networkTypes []
wg.Add(1)
go func() {
defer wg.Done()
conn, err := listenUDPInPortRange(a.net, a.log, int(a.portmax), int(a.portmin), network, &net.UDPAddr{IP: nil, Port: 0})
if err != nil {
a.log.Warnf("Failed to listen %s: %v\n", network, err)
@@ -291,6 +292,7 @@ func (a *Agent) gatherCandidatesSrflx(ctx context.Context, urls []*URL, networkT
wg.Add(1)
go func(url URL, network string) {
defer wg.Done()
hostPort := fmt.Sprintf("%s:%d", url.Host, url.Port)
serverAddr, err := a.net.ResolveUDPAddr(network, hostPort)
if err != nil {
+82 -36
View File
@@ -1,9 +1,14 @@
package ice
import (
"errors"
"fmt"
"io"
"net"
"os"
"strings"
"sync"
"time"
"github.com/pion/logging"
)
@@ -11,8 +16,9 @@ import (
// UDPMux allows multiple connections to go over a single UDP port
type UDPMux interface {
io.Closer
GetConnByUfrag(ufrag string) (net.PacketConn, error)
GetConn(ufrag, network string) (net.PacketConn, error)
RemoveConnByUfrag(ufrag string)
Start(port int) error
}
// UDPMuxDefault is an implementation of the interface
@@ -25,7 +31,7 @@ type UDPMuxDefault struct {
closedChan chan struct{}
closeOnce sync.Once
// conns is a map of all udpMuxedConn indexed by ufrag
// conns is a map of all udpMuxedConn indexed by ufrag|network|candidateType
conns map[string]*udpMuxedConn
// buffer pool to recycle buffers for incoming packets
@@ -54,7 +60,7 @@ func NewUDPMuxDefault(params UDPMuxParams) *UDPMuxDefault {
closedChan: make(chan struct{}, 1),
pool: &sync.Pool{
New: func() interface{} {
return make([]byte, receiveMTU)
return newBufferHolder(receiveMTU)
},
},
}
@@ -83,13 +89,15 @@ func (m *UDPMuxDefault) LocalAddr() net.Addr {
return m.listenAddr
}
// GetConnByUfrag returns a PacketConn given the connection's ufrag.
// GetConn returns a PacketConn given the connection's ufrag and network
// creates the connection if an existing one can't be found
func (m *UDPMuxDefault) GetConnByUfrag(ufrag string) (net.PacketConn, error) {
func (m *UDPMuxDefault) GetConn(ufrag, network string) (net.PacketConn, error) {
if m.udpConn == nil {
return nil, ErrMuxNotStarted
}
key := fmt.Sprintf("%s|%s", ufrag, network)
m.mu.Lock()
defer m.mu.Unlock()
@@ -97,37 +105,44 @@ func (m *UDPMuxDefault) GetConnByUfrag(ufrag string) (net.PacketConn, error) {
return nil, io.ErrClosedPipe
}
if c, ok := m.conns[ufrag]; ok {
if c, ok := m.conns[key]; ok {
return c, nil
}
c := m.createMuxedConn()
go func() {
<-c.CloseChannel()
m.RemoveConnByUfrag(ufrag)
print("muxed connection closed, removing key ", key, "\n")
m.removeConn(key)
}()
m.conns[ufrag] = c
m.conns[key] = c
return c, nil
}
// RemoveConnByUfrag stops and removes the muxed packet connection
func (m *UDPMuxDefault) RemoveConnByUfrag(ufrag string) {
// get addresses to remove
m.mu.Lock()
c := m.conns[ufrag]
delete(m.conns, ufrag)
removedConns := make([]*udpMuxedConn, 0)
for key := range m.conns {
if !strings.HasPrefix(key, ufrag) {
continue
}
c := m.conns[key]
delete(m.conns, key)
if c != nil {
removedConns = append(removedConns, c)
}
}
// keep lock section small to avoid deadlock with conn lock
m.mu.Unlock()
if c == nil {
return
}
addresses := c.getAddresses()
for _, addr := range addresses {
m.mappingChan <- connMap{
address: addr,
conn: nil,
for _, c := range removedConns {
addresses := c.getAddresses()
for _, addr := range addresses {
m.mappingChan <- connMap{
address: addr,
conn: nil,
}
}
}
}
@@ -161,12 +176,31 @@ func (m *UDPMuxDefault) Close() error {
return err
}
func (m *UDPMuxDefault) removeConn(key string) {
m.mu.Lock()
c := m.conns[key]
delete(m.conns, key)
// keep lock section small to avoid deadlock with conn lock
m.mu.Unlock()
if c == nil {
return
}
addresses := c.getAddresses()
for _, addr := range addresses {
m.mappingChan <- connMap{
address: addr,
conn: nil,
}
}
}
func (m *UDPMuxDefault) writeTo(buf []byte, raddr net.Addr) (n int, err error) {
return m.udpConn.WriteTo(buf, raddr)
}
func (m *UDPMuxDefault) doneWithBuffer(buf []byte) {
//nolint
func (m *UDPMuxDefault) doneWithBuffer(buf *bufferHolder) {
m.pool.Put(buf)
}
@@ -200,33 +234,35 @@ func (m *UDPMuxDefault) connWorker() {
_ = m.Close()
}()
for {
buffer := m.pool.Get().([]byte)
n, addr, err := m.udpConn.ReadFrom(buffer)
if err == io.EOF {
return
} else if err != nil {
logger.Errorf("could not read udp packet: %v", err)
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
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)
}
return
}
// process any mapping changes
m.applyMappingChanges(remoteMap)
// look up forward destination
addrStr := addr.String()
c := remoteMap[addrStr]
if c == nil {
//nolint
m.pool.Put(buffer)
m.doneWithBuffer(buffer)
// ignore packets that we don't know where to route to
continue
}
err = c.writePacket(muxedPacket{
Data: buffer,
Size: n,
RAddr: addr,
Buffer: buffer,
Size: n,
RAddr: addr,
})
if err != nil {
logger.Errorf("could not write packet: %v", err)
@@ -249,3 +285,13 @@ func (m *UDPMuxDefault) applyMappingChanges(remoteMap map[string]*udpMuxedConn)
}
}
}
type bufferHolder struct {
buffer []byte
}
func newBufferHolder(size int) *bufferHolder {
return &bufferHolder{
buffer: make([]byte, size),
}
}
+2 -2
View File
@@ -52,12 +52,12 @@ func TestUDPMux(t *testing.T) {
require.NoError(t, udpMux.Close())
// can't create more connections
_, err = udpMux.GetConnByUfrag("failufrag")
_, err = udpMux.GetConn("failufrag", "udp")
require.Error(t, err)
}
func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string) {
pktConn, err := udpMux.GetConnByUfrag(ufrag)
pktConn, err := udpMux.GetConn(ufrag, udp)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
+5 -5
View File
@@ -17,9 +17,9 @@ type udpMuxedConnParams struct {
}
type muxedPacket struct {
Data []byte
RAddr net.Addr
Size int
Buffer *bufferHolder
RAddr net.Addr
Size int
}
// udpMuxedConn represents a logical packet conn for a single remote as identified by ufrag
@@ -56,8 +56,8 @@ func (c *udpMuxedConn) ReadFrom(b []byte) (n int, raddr net.Addr, err error) {
return 0, pkt.RAddr, io.ErrShortBuffer
}
copy(b, pkt.Data[:pkt.Size])
c.params.Mux.doneWithBuffer(pkt.Data)
copy(b, pkt.Buffer.buffer[:pkt.Size])
c.params.Mux.doneWithBuffer(pkt.Buffer)
return pkt.Size, pkt.RAddr, err
}