mirror of
https://github.com/netbirdio/ice.git
synced 2026-05-22 17:10:58 -07:00
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user