Fix muxed conn shared by candidate

When we have multiple host candidates and mux to a
single port, if these candidates share a same conn
(either tcp or udp), they might read other's msg
cause ice connection cost long time or failed.
This commit is contained in:
cnderrauber
2022-09-26 09:26:23 +08:00
committed by cnderrauber
parent 5475fdb796
commit 93980395c8
20 changed files with 443 additions and 135 deletions
+1
View File
@@ -27,6 +27,7 @@ Jerko Steiner <jerko.steiner@gmail.com>
JooYoung <qkdlql@naver.com>
Juliusz Chroboczek <jch@irif.fr>
Kacper Bąk <56700396+53jk1@users.noreply.github.com>
Kevin Caffrey <kcaffrey@gmail.com>
Konstantin Itskov <konstantin.itskov@kovits.com>
korymiller1489 <kmiller@unwiredrevolution.com>
Kyle Carberry <kyle@carberry.com>
+5 -3
View File
@@ -1610,9 +1610,11 @@ func TestRunTaskInSelectedCandidatePairChangeCallback(t *testing.T) {
isComplete := make(chan interface{})
isTested := make(chan interface{})
if err = aAgent.OnSelectedCandidatePairChange(func(Candidate, Candidate) {
_, _, errCred := aAgent.GetLocalUserCredentials()
assert.NoError(t, errCred)
close(isTested)
go func() {
_, _, errCred := aAgent.GetLocalUserCredentials()
assert.NoError(t, errCred)
close(isTested)
}()
}); err != nil {
t.Error(err)
}
+1 -1
View File
@@ -23,7 +23,7 @@ func TestMuxAgent(t *testing.T) {
const muxPort = 7686
c, err := net.ListenUDP(udp, &net.UDPAddr{
c, err := net.ListenUDP("udp4", &net.UDPAddr{
Port: muxPort,
})
+5 -4
View File
@@ -186,7 +186,7 @@ func (a *Agent) gatherCandidatesLocal(ctx context.Context, networkTypes []Networ
var muxConns []net.PacketConn
if multi, ok := a.tcpMux.(AllConnsGetter); ok {
a.log.Debugf("GetAllConns by ufrag: %s", a.localUfrag)
muxConns, err = multi.GetAllConns(a.localUfrag, mappedIP.To4() == nil)
muxConns, err = multi.GetAllConns(a.localUfrag, mappedIP.To4() == nil, ip)
if err != nil {
if !errors.Is(err, ErrTCPMuxNotInitialized) {
a.log.Warnf("error getting all tcp conns by ufrag: %s %s %s", network, ip, a.localUfrag)
@@ -195,7 +195,7 @@ func (a *Agent) gatherCandidatesLocal(ctx context.Context, networkTypes []Networ
}
} else {
a.log.Debugf("GetConn by ufrag: %s", a.localUfrag)
conn, err := a.tcpMux.GetConnByUfrag(a.localUfrag, mappedIP.To4() == nil)
conn, err := a.tcpMux.GetConnByUfrag(a.localUfrag, mappedIP.To4() == nil, ip)
if err != nil {
if !errors.Is(err, ErrTCPMuxNotInitialized) {
a.log.Warnf("error getting tcp conn by ufrag: %s %s %s", network, ip, a.localUfrag)
@@ -282,6 +282,7 @@ func (a *Agent) gatherCandidatesLocalUDPMux(ctx context.Context) error { //nolin
}
for _, candidateIP := range localIPs {
localIP := candidateIP
if a.extIPMapper != nil && a.extIPMapper.candidateType == CandidateTypeHost {
if mappedIP, innerErr := a.extIPMapper.findExternalIP(candidateIP.String()); innerErr != nil {
a.log.Warnf("1:1 NAT mapping is enabled but no external IP is found for %s", candidateIP.String())
@@ -293,7 +294,7 @@ func (a *Agent) gatherCandidatesLocalUDPMux(ctx context.Context) error { //nolin
var conns []net.PacketConn
if multi, ok := a.udpMux.(AllConnsGetter); ok {
conns, err = multi.GetAllConns(a.localUfrag, candidateIP.To4() == nil)
conns, err = multi.GetAllConns(a.localUfrag, candidateIP.To4() == nil, localIP)
if err != nil {
return err
}
@@ -302,7 +303,7 @@ func (a *Agent) gatherCandidatesLocalUDPMux(ctx context.Context) error { //nolin
continue
}
} else {
conn, err := a.udpMux.GetConn(a.localUfrag, candidateIP.To4() == nil)
conn, err := a.udpMux.GetConn(a.localUfrag, candidateIP.To4() == nil, localIP)
if err != nil {
return err
}
+4
View File
@@ -510,6 +510,10 @@ func TestMultiUDPMuxUsage(t *testing.T) {
udpMuxInstances = append(udpMuxInstances, NewUDPMuxDefault(UDPMuxParams{
UDPConn: conn,
}))
idx := i
defer func() {
_ = udpMuxInstances[idx].Close()
}()
}
a, err := NewAgent(&AgentConfig{
+2 -1
View File
@@ -252,7 +252,8 @@ func (s *controlledSelector) HandleBindingRequest(m *stun.Message, local, remote
// previously sent by this pair produced a successful response and
// generated a valid pair (Section 7.2.5.3.2). The agent sets the
// nominated flag value of the valid pair to true.
if selectedPair := s.agent.getSelectedPair(); selectedPair == nil || selectedPair.priority() < p.priority() {
if selectedPair := s.agent.getSelectedPair(); selectedPair == nil ||
(selectedPair != p && selectedPair.priority() <= p.priority()) {
s.agent.setSelectedPair(p)
} else if selectedPair != p {
s.log.Tracef("ignore nominate new pair %s, already nominated pair %s", p, selectedPair)
+110 -31
View File
@@ -2,6 +2,7 @@ package ice
import (
"encoding/binary"
"errors"
"io"
"net"
"strings"
@@ -11,6 +12,9 @@ import (
"github.com/pion/stun"
)
// ErrGetTransportAddress can't convert net.Addr to underlying type (UDPAddr or TCPAddr).
var ErrGetTransportAddress = errors.New("failed to get local transport address")
// TCPMux is allows grouping multiple TCP net.Conns and using them like UDP
// net.PacketConns. The main implementation of this is TCPMuxDefault, and this
// interface exists to:
@@ -19,7 +23,7 @@ import (
// 2. allow mocking in tests.
type TCPMux interface {
io.Closer
GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketConn, error)
GetConnByUfrag(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error)
RemoveConnByUfrag(ufrag string)
}
@@ -36,21 +40,23 @@ func (m *invalidTCPMux) Close() error {
}
// GetConnByUfrag implements TCPMux interface.
func (m *invalidTCPMux) GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketConn, error) {
func (m *invalidTCPMux) GetConnByUfrag(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error) {
return nil, ErrTCPMuxNotInitialized
}
// RemoveConnByUfrag implements TCPMux interface.
func (m *invalidTCPMux) RemoveConnByUfrag(ufrag string) {}
type ipAddr string
// TCPMuxDefault muxes TCP net.Conns into net.PacketConns and groups them by
// Ufrag. It is a default implementation of TCPMux interface.
type TCPMuxDefault struct {
params *TCPMuxParams
closed bool
// connsIPv4 and connsIPv6 are maps of all tcpPacketConns indexed by ufrag
connsIPv4, connsIPv6 map[string]*tcpPacketConn
// connsIPv4 and connsIPv6 are maps of all tcpPacketConns indexed by ufrag and local address
connsIPv4, connsIPv6 map[string]map[ipAddr]*tcpPacketConn
mu sync.Mutex
wg sync.WaitGroup
@@ -77,8 +83,8 @@ func NewTCPMuxDefault(params TCPMuxParams) *TCPMuxDefault {
m := &TCPMuxDefault{
params: &params,
connsIPv4: map[string]*tcpPacketConn{},
connsIPv6: map[string]*tcpPacketConn{},
connsIPv4: map[string]map[ipAddr]*tcpPacketConn{},
connsIPv6: map[string]map[ipAddr]*tcpPacketConn{},
}
m.wg.Add(1)
@@ -115,7 +121,7 @@ func (m *TCPMuxDefault) LocalAddr() net.Addr {
}
// GetConnByUfrag retrieves an existing or creates a new net.PacketConn.
func (m *TCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketConn, error) {
func (m *TCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error) {
m.mu.Lock()
defer m.mu.Unlock()
@@ -123,35 +129,50 @@ func (m *TCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketCon
return nil, io.ErrClosedPipe
}
if conn, ok := m.getConn(ufrag, isIPv6); ok {
if conn, ok := m.getConn(ufrag, isIPv6, local); ok {
return conn, nil
}
return m.createConn(ufrag, m.LocalAddr(), isIPv6), nil
return m.createConn(ufrag, isIPv6, local)
}
func (m *TCPMuxDefault) createConn(ufrag string, localAddr net.Addr, isIPv6 bool) *tcpPacketConn {
func (m *TCPMuxDefault) createConn(ufrag string, isIPv6 bool, local net.IP) (*tcpPacketConn, error) {
addr, ok := m.LocalAddr().(*net.TCPAddr)
if !ok {
return nil, ErrGetTransportAddress
}
localAddr := *addr
localAddr.IP = local
conn := newTCPPacketConn(tcpPacketParams{
ReadBuffer: m.params.ReadBufferSize,
WriteBuffer: m.params.WriteBufferSize,
LocalAddr: localAddr,
LocalAddr: &localAddr,
Logger: m.params.Logger,
})
var conns map[ipAddr]*tcpPacketConn
if isIPv6 {
m.connsIPv6[ufrag] = conn
if conns, ok = m.connsIPv6[ufrag]; !ok {
conns = make(map[ipAddr]*tcpPacketConn)
m.connsIPv6[ufrag] = conns
}
} else {
m.connsIPv4[ufrag] = conn
if conns, ok = m.connsIPv4[ufrag]; !ok {
conns = make(map[ipAddr]*tcpPacketConn)
m.connsIPv4[ufrag] = conns
}
}
conns[ipAddr(local.String())] = conn
m.wg.Add(1)
go func() {
defer m.wg.Done()
<-conn.CloseChannel()
m.RemoveConnByUfrag(ufrag)
m.removeConnByUfragAndLocalHost(ufrag, local)
}()
return conn
return conn, nil
}
func (m *TCPMuxDefault) closeAndLogError(closer io.Closer) {
@@ -214,9 +235,21 @@ func (m *TCPMuxDefault) handleConn(conn net.Conn) {
}
isIPv6 := net.ParseIP(host).To4() == nil
packetConn, ok := m.getConn(ufrag, isIPv6)
localAddr, ok := conn.LocalAddr().(*net.TCPAddr)
if !ok {
packetConn = m.createConn(ufrag, conn.LocalAddr(), isIPv6)
m.closeAndLogError(conn)
m.params.Logger.Warnf("Failed to get local tcp address in STUN message from %s to %s", conn.RemoteAddr(), conn.LocalAddr())
return
}
packetConn, ok := m.getConn(ufrag, isIPv6, localAddr.IP)
if !ok {
packetConn, err = m.createConn(ufrag, isIPv6, localAddr.IP)
if err != nil {
m.closeAndLogError(conn)
m.params.Logger.Warnf("Failed to create packetConn for STUN message from %s to %s", conn.RemoteAddr(), conn.LocalAddr())
return
}
}
if err := packetConn.AddConn(conn, buf); err != nil {
@@ -231,15 +264,19 @@ func (m *TCPMuxDefault) Close() error {
m.mu.Lock()
m.closed = true
for _, conn := range m.connsIPv4 {
m.closeAndLogError(conn)
for _, conns := range m.connsIPv4 {
for _, conn := range conns {
m.closeAndLogError(conn)
}
}
for _, conn := range m.connsIPv6 {
m.closeAndLogError(conn)
for _, conns := range m.connsIPv6 {
for _, conn := range conns {
m.closeAndLogError(conn)
}
}
m.connsIPv4 = map[string]*tcpPacketConn{}
m.connsIPv6 = map[string]*tcpPacketConn{}
m.connsIPv4 = map[string]map[ipAddr]*tcpPacketConn{}
m.connsIPv6 = map[string]map[ipAddr]*tcpPacketConn{}
err := m.params.Listener.Close()
@@ -252,17 +289,55 @@ func (m *TCPMuxDefault) Close() error {
// RemoveConnByUfrag closes and removes a net.PacketConn by Ufrag.
func (m *TCPMuxDefault) RemoveConnByUfrag(ufrag string) {
removedConns := make([]*tcpPacketConn, 0, 2)
removedConns := make([]*tcpPacketConn, 0, 4)
// Keep lock section small to avoid deadlock with conn lock
m.mu.Lock()
if conn, ok := m.connsIPv4[ufrag]; ok {
if conns, ok := m.connsIPv4[ufrag]; ok {
delete(m.connsIPv4, ufrag)
removedConns = append(removedConns, conn)
for _, conn := range conns {
removedConns = append(removedConns, conn)
}
}
if conn, ok := m.connsIPv6[ufrag]; ok {
if conns, ok := m.connsIPv6[ufrag]; ok {
delete(m.connsIPv6, ufrag)
removedConns = append(removedConns, conn)
for _, conn := range conns {
removedConns = append(removedConns, conn)
}
}
m.mu.Unlock()
// Close the connections outside the critical section to avoid
// deadlocking TCP mux if (*tcpPacketConn).Close() blocks.
for _, conn := range removedConns {
m.closeAndLogError(conn)
}
}
func (m *TCPMuxDefault) removeConnByUfragAndLocalHost(ufrag string, local net.IP) {
removedConns := make([]*tcpPacketConn, 0, 4)
localIP := ipAddr(local.String())
// Keep lock section small to avoid deadlock with conn lock
m.mu.Lock()
if conns, ok := m.connsIPv4[ufrag]; ok {
if conn, ok := conns[localIP]; ok {
delete(conns, localIP)
if len(conns) == 0 {
delete(m.connsIPv4, ufrag)
}
removedConns = append(removedConns, conn)
}
}
if conns, ok := m.connsIPv6[ufrag]; ok {
if conn, ok := conns[localIP]; ok {
delete(conns, localIP)
if len(conns) == 0 {
delete(m.connsIPv6, ufrag)
}
removedConns = append(removedConns, conn)
}
}
m.mu.Unlock()
@@ -273,11 +348,15 @@ func (m *TCPMuxDefault) RemoveConnByUfrag(ufrag string) {
}
}
func (m *TCPMuxDefault) getConn(ufrag string, isIPv6 bool) (val *tcpPacketConn, ok bool) {
func (m *TCPMuxDefault) getConn(ufrag string, isIPv6 bool, local net.IP) (val *tcpPacketConn, ok bool) {
var conns map[ipAddr]*tcpPacketConn
if isIPv6 {
val, ok = m.connsIPv6[ufrag]
conns, ok = m.connsIPv6[ufrag]
} else {
val, ok = m.connsIPv4[ufrag]
conns, ok = m.connsIPv4[ufrag]
}
if conns != nil {
val, ok = conns[ipAddr(local.String())]
}
return
+5 -5
View File
@@ -20,17 +20,17 @@ func NewMultiTCPMuxDefault(muxs ...TCPMux) *MultiTCPMuxDefault {
}
}
// GetConnByUfrag returns a PacketConn given the connection's ufrag and network
// GetConnByUfrag returns a PacketConn given the connection's ufrag, network and local address
// creates the connection if an existing one can't be found. This, unlike
// GetAllConns, will only return a single PacketConn from the first mux that was
// passed in to NewMultiTCPMuxDefault.
func (m *MultiTCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketConn, error) {
func (m *MultiTCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error) {
// NOTE: We always use the first element here in order to maintain the
// behavior of using an existing connection if one exists.
if len(m.muxs) == 0 {
return nil, errNoTCPMuxAvailable
}
return m.muxs[0].GetConnByUfrag(ufrag, isIPv6)
return m.muxs[0].GetConnByUfrag(ufrag, isIPv6, local)
}
// RemoveConnByUfrag stops and removes the muxed packet connection
@@ -42,14 +42,14 @@ func (m *MultiTCPMuxDefault) RemoveConnByUfrag(ufrag string) {
}
// GetAllConns returns a PacketConn for each underlying TCPMux
func (m *MultiTCPMuxDefault) GetAllConns(ufrag string, isIPv6 bool) ([]net.PacketConn, error) {
func (m *MultiTCPMuxDefault) GetAllConns(ufrag string, isIPv6 bool, local net.IP) ([]net.PacketConn, error) {
if len(m.muxs) == 0 {
// Make sure that we either return at least one connection or an error.
return nil, errNoTCPMuxAvailable
}
var conns []net.PacketConn
for _, mux := range m.muxs {
conn, err := mux.GetConnByUfrag(ufrag, isIPv6)
conn, err := mux.GetConnByUfrag(ufrag, isIPv6, local)
if err != nil {
// For now, this implementation is all or none.
return nil, err
+3 -3
View File
@@ -53,7 +53,7 @@ func TestMultiTCPMux_Recv(t *testing.T) {
_ = multiMux.Close()
}()
pktConns, err := multiMux.GetAllConns("myufrag", false)
pktConns, err := multiMux.GetAllConns("myufrag", false, net.IP{127, 0, 0, 1})
require.NoError(t, err, "error retrieving muxed connection for ufrag")
for _, pktConn := range pktConns {
@@ -117,12 +117,12 @@ func TestMultiTCPMux_NoDeadlockWhenClosingUnusedPacketConn(t *testing.T) {
}
muxMulti := NewMultiTCPMuxDefault(tcpMuxInstances...)
_, err := muxMulti.GetAllConns("test", false)
_, err := muxMulti.GetAllConns("test", false, net.IP{127, 0, 0, 1})
require.NoError(t, err, "error getting conn by ufrag")
require.NoError(t, muxMulti.Close(), "error closing tcpMux")
conn, err := muxMulti.GetAllConns("test", false)
conn, err := muxMulti.GetAllConns("test", false, net.IP{127, 0, 0, 1})
assert.Nil(t, conn, "should receive nil because mux is closed")
assert.Equal(t, io.ErrClosedPipe, err, "should receive error because mux is closed")
}
+3 -3
View File
@@ -62,7 +62,7 @@ func TestTCPMux_Recv(t *testing.T) {
n, err := writeStreamingPacket(conn, msg.Raw)
require.NoError(t, err, "error writing tcp stun packet")
pktConn, err := tcpMux.GetConnByUfrag("myufrag", false)
pktConn, err := tcpMux.GetConnByUfrag("myufrag", false, listener.Addr().(*net.TCPAddr).IP)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
@@ -108,12 +108,12 @@ func TestTCPMux_NoDeadlockWhenClosingUnusedPacketConn(t *testing.T) {
ReadBufferSize: 20,
})
_, err = tcpMux.GetConnByUfrag("test", false)
_, err = tcpMux.GetConnByUfrag("test", false, listener.Addr().(*net.TCPAddr).IP)
require.NoError(t, err, "error getting conn by ufrag")
require.NoError(t, tcpMux.Close(), "error closing tcpMux")
conn, err := tcpMux.GetConnByUfrag("test", false)
conn, err := tcpMux.GetConnByUfrag("test", false, listener.Addr().(*net.TCPAddr).IP)
assert.Nil(t, conn, "should receive nil because mux is closed")
assert.Equal(t, io.ErrClosedPipe, err, "should receive error because mux is closed")
}
+1 -1
View File
@@ -111,7 +111,7 @@ func newTCPPacketConn(params tcpPacketParams) *tcpPacketConn {
}
func (t *tcpPacketConn) AddConn(conn net.Conn, firstPacketData []byte) error {
t.params.Logger.Infof("AddConn: %s %s", conn.RemoteAddr().Network(), conn.RemoteAddr())
t.params.Logger.Infof("AddConn: %s remote %s to local %s", conn.RemoteAddr().Network(), conn.RemoteAddr(), conn.LocalAddr())
t.mu.Lock()
defer t.mu.Unlock()
+185 -46
View File
@@ -15,7 +15,7 @@ import (
// UDPMux allows multiple connections to go over a single UDP port
type UDPMux interface {
io.Closer
GetConn(ufrag string, isIPv6 bool) (net.PacketConn, error)
GetConn(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error)
RemoveConnByUfrag(ufrag string)
}
@@ -27,10 +27,12 @@ type UDPMuxDefault struct {
closeOnce sync.Once
// connsIPv4 and connsIPv6 are maps of all udpMuxedConn indexed by ufrag|network|candidateType
connsIPv4, connsIPv6 map[string]*udpMuxedConn
connsIPv4, connsIPv6 map[string]map[ipAddr]*udpMuxedConn
addressMapMu sync.RWMutex
addressMap map[string]*udpMuxedConn
// remote address (ip:port) -> (localip -> udpMuxedConn)
addressMap map[string]map[ipAddr]*udpMuxedConn
// buffer pool to recycle buffers for net.UDPAddr encodes/decodes
pool *sync.Pool
@@ -40,10 +42,24 @@ type UDPMuxDefault struct {
const maxAddrSize = 512
// UDPMuxConn is a udp PacketConn with ReadMsgUDP and File method
// to retrieve the destination local address of the received packet
type UDPMuxConn interface {
net.PacketConn
// ReadMsgUdp used to get destination address when received a udp packet
ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net.UDPAddr, err error)
// File returns a copy of the underlying os.File.
// It is the caller's responsibility to close f when finished.
// Closing c does not affect f, and closing f does not affect c.
File() (f *os.File, err error)
}
// UDPMuxParams are parameters for UDPMux.
type UDPMuxParams struct {
Logger logging.LeveledLogger
UDPConn net.PacketConn
UDPConn UDPMuxConn
}
// NewUDPMuxDefault creates an implementation of UDPMux
@@ -53,10 +69,10 @@ func NewUDPMuxDefault(params UDPMuxParams) *UDPMuxDefault {
}
m := &UDPMuxDefault{
addressMap: map[string]*udpMuxedConn{},
addressMap: make(map[string]map[ipAddr]*udpMuxedConn),
params: params,
connsIPv4: make(map[string]*udpMuxedConn),
connsIPv6: make(map[string]*udpMuxedConn),
connsIPv4: make(map[string]map[ipAddr]*udpMuxedConn),
connsIPv6: make(map[string]map[ipAddr]*udpMuxedConn),
closedChan: make(chan struct{}, 1),
pool: &sync.Pool{
New: func() interface{} {
@@ -78,7 +94,7 @@ func (m *UDPMuxDefault) LocalAddr() net.Addr {
// 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) GetConn(ufrag string, isIPv6 bool) (net.PacketConn, error) {
func (m *UDPMuxDefault) GetConn(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error) {
m.mu.Lock()
defer m.mu.Unlock()
@@ -86,38 +102,56 @@ func (m *UDPMuxDefault) GetConn(ufrag string, isIPv6 bool) (net.PacketConn, erro
return nil, io.ErrClosedPipe
}
if conn, ok := m.getConn(ufrag, isIPv6); ok {
if conn, ok := m.getConn(ufrag, isIPv6, local); ok {
return conn, nil
}
c := m.createMuxedConn(ufrag)
c, err := m.createMuxedConn(ufrag, local)
if err != nil {
return nil, err
}
go func() {
<-c.CloseChannel()
m.RemoveConnByUfrag(ufrag)
m.removeConnByUfragAndLocalHost(ufrag, local)
}()
var (
conns map[ipAddr]*udpMuxedConn
ok bool
)
if isIPv6 {
m.connsIPv6[ufrag] = c
if conns, ok = m.connsIPv6[ufrag]; !ok {
conns = make(map[ipAddr]*udpMuxedConn)
m.connsIPv6[ufrag] = conns
}
} else {
m.connsIPv4[ufrag] = c
if conns, ok = m.connsIPv4[ufrag]; !ok {
conns = make(map[ipAddr]*udpMuxedConn)
m.connsIPv4[ufrag] = conns
}
}
conns[ipAddr(local.String())] = c
return c, nil
}
// RemoveConnByUfrag stops and removes the muxed packet connection
func (m *UDPMuxDefault) RemoveConnByUfrag(ufrag string) {
removedConns := make([]*udpMuxedConn, 0, 2)
removedConns := make([]*udpMuxedConn, 0, 4)
// Keep lock section small to avoid deadlock with conn lock
m.mu.Lock()
if c, ok := m.connsIPv4[ufrag]; ok {
if conns, ok := m.connsIPv4[ufrag]; ok {
delete(m.connsIPv4, ufrag)
removedConns = append(removedConns, c)
for _, c := range conns {
removedConns = append(removedConns, c)
}
}
if c, ok := m.connsIPv6[ufrag]; ok {
if conns, ok := m.connsIPv6[ufrag]; ok {
delete(m.connsIPv6, ufrag)
removedConns = append(removedConns, c)
for _, c := range conns {
removedConns = append(removedConns, c)
}
}
m.mu.Unlock()
@@ -132,7 +166,59 @@ func (m *UDPMuxDefault) RemoveConnByUfrag(ufrag string) {
for _, c := range removedConns {
addresses := c.getAddresses()
for _, addr := range addresses {
delete(m.addressMap, addr)
if conns, ok := m.addressMap[addr]; ok {
delete(conns, ipAddr(c.params.LocalIP.String()))
if len(conns) == 0 {
delete(m.addressMap, addr)
}
}
}
}
}
func (m *UDPMuxDefault) removeConnByUfragAndLocalHost(ufrag string, local net.IP) {
removedConns := make([]*udpMuxedConn, 0, 4)
localIP := ipAddr(local.String())
// Keep lock section small to avoid deadlock with conn lock
m.mu.Lock()
if conns, ok := m.connsIPv4[ufrag]; ok {
if conn, ok := conns[localIP]; ok {
delete(conns, localIP)
if len(conns) == 0 {
delete(m.connsIPv4, ufrag)
}
removedConns = append(removedConns, conn)
}
}
if conns, ok := m.connsIPv6[ufrag]; ok {
if conn, ok := conns[localIP]; ok {
delete(conns, localIP)
if len(conns) == 0 {
delete(m.connsIPv6, ufrag)
}
removedConns = append(removedConns, conn)
}
}
m.mu.Unlock()
if len(removedConns) == 0 {
// No need to lock if no connection was found
return
}
m.addressMapMu.Lock()
defer m.addressMapMu.Unlock()
for _, c := range removedConns {
addresses := c.getAddresses()
for _, addr := range addresses {
if conns, ok := m.addressMap[addr]; ok {
delete(conns, ipAddr(c.params.LocalIP.String()))
if len(conns) == 0 {
delete(m.addressMap, addr)
}
}
}
}
}
@@ -154,17 +240,40 @@ func (m *UDPMuxDefault) Close() error {
m.mu.Lock()
defer m.mu.Unlock()
for _, c := range m.connsIPv4 {
_ = c.Close()
for _, conns := range m.connsIPv4 {
for _, c := range conns {
_ = c.Close()
}
}
for _, c := range m.connsIPv6 {
_ = c.Close()
for _, conns := range m.connsIPv6 {
for _, c := range conns {
_ = c.Close()
}
}
m.connsIPv4 = make(map[string]*udpMuxedConn)
m.connsIPv6 = make(map[string]*udpMuxedConn)
m.connsIPv4 = make(map[string]map[ipAddr]*udpMuxedConn)
m.connsIPv6 = make(map[string]map[ipAddr]*udpMuxedConn)
// ReadMsgUDP will block until something is received, otherwise it will block forever
// and the Conn's Close method too. So send a packet to wake it for exit.
close(m.closedChan)
closeConn, errConn := net.DialUDP("udp", nil, m.params.UDPConn.LocalAddr().(*net.UDPAddr))
// i386 doesn't support dial local ipv6 address
if errConn != nil && strings.Contains(errConn.Error(), "dial udp [::]:") &&
strings.Contains(errConn.Error(), "connect: cannot assign requested address") {
closeConn, errConn = net.DialUDP("udp4", nil, &net.UDPAddr{Port: m.params.UDPConn.LocalAddr().(*net.UDPAddr).Port})
}
if errConn != nil {
m.params.Logger.Errorf("Failed to open close notify socket, %v", errConn)
} else {
defer func() {
_ = closeConn.Close()
}()
_, errConn = closeConn.Write([]byte("close"))
if errConn != nil {
m.params.Logger.Errorf("Failed to send close notify msg, %v", errConn)
}
}
})
return err
}
@@ -181,36 +290,58 @@ func (m *UDPMuxDefault) registerConnForAddress(conn *udpMuxedConn, addr string)
m.addressMapMu.Lock()
defer m.addressMapMu.Unlock()
existing, ok := m.addressMap[addr]
conns, ok := m.addressMap[addr]
if ok {
existing.removeAddress(addr)
existing, ok := conns[ipAddr(conn.params.LocalIP.String())]
if ok {
existing.removeAddress(addr)
}
} else {
conns = make(map[ipAddr]*udpMuxedConn)
m.addressMap[addr] = conns
}
m.addressMap[addr] = conn
conns[ipAddr(conn.params.LocalIP.String())] = conn
m.params.Logger.Debugf("Registered %s for %s", addr, conn.params.Key)
m.params.Logger.Debugf("Registered %s for %s, local %s", addr, conn.params.Key, conn.params.LocalIP.String())
}
func (m *UDPMuxDefault) createMuxedConn(key string) *udpMuxedConn {
func (m *UDPMuxDefault) createMuxedConn(key string, local net.IP) (*udpMuxedConn, error) {
m.params.Logger.Debugf("Creating new muxed connection, key:%s local:%s ", key, local.String())
addr, ok := m.LocalAddr().(*net.UDPAddr)
if !ok {
return nil, ErrGetTransportAddress
}
localAddr := *addr
localAddr.IP = local
c := newUDPMuxedConn(&udpMuxedConnParams{
Mux: m,
Key: key,
AddrPool: m.pool,
LocalAddr: m.LocalAddr(),
LocalAddr: &localAddr,
LocalIP: local,
Logger: m.params.Logger,
})
return c
return c, nil
}
func (m *UDPMuxDefault) connWorker() {
func (m *UDPMuxDefault) connWorker() { //nolint:gocognit
logger := m.params.Logger
defer func() {
_ = m.Close()
}()
localUDPAddr, _ := m.LocalAddr().(*net.UDPAddr)
buf := make([]byte, receiveMTU)
file, _ := m.params.UDPConn.File()
setUDPSocketOptionsForLocalAddr(file.Fd(), m.params.Logger)
_ = file.Close()
oob := make([]byte, receiveMTU)
for {
n, addr, err := m.params.UDPConn.ReadFrom(buf)
localHost := localUDPAddr.IP
n, oobn, _, addr, err := m.params.UDPConn.ReadMsgUDP(buf, oob)
if m.IsClosed() {
return
} else if err != nil {
@@ -223,15 +354,19 @@ func (m *UDPMuxDefault) connWorker() {
return
}
udpAddr, ok := addr.(*net.UDPAddr)
if !ok {
logger.Errorf("underlying PacketConn did not return a UDPAddr")
return
// get destination local addr from received packet
if oobIP, addrErr := getLocalAddrFromOob(oob[:oobn]); addrErr == nil {
localHost = oobIP
} else {
m.params.Logger.Warnf("could not get local addr from oob: %v, remote %s", addrErr, addr)
}
// If we have already seen this address dispatch to the appropriate destination
var destinationConn *udpMuxedConn
m.addressMapMu.Lock()
destinationConn := m.addressMap[addr.String()]
if conns, ok := m.addressMap[addr.String()]; ok {
destinationConn = conns[ipAddr(localHost.String())]
}
m.addressMapMu.Unlock()
// If we haven't seen this address before but is a STUN packet lookup by ufrag
@@ -252,29 +387,33 @@ func (m *UDPMuxDefault) connWorker() {
}
ufrag := strings.Split(string(attr), ":")[0]
isIPv6 := udpAddr.IP.To4() == nil
isIPv6 := addr.IP.To4() == nil
m.mu.Lock()
destinationConn, _ = m.getConn(ufrag, isIPv6)
destinationConn, _ = m.getConn(ufrag, isIPv6, localHost)
m.mu.Unlock()
}
if destinationConn == nil {
m.params.Logger.Tracef("dropping packet from %s, addr: %s", udpAddr.String(), addr.String())
m.params.Logger.Tracef("dropping packet from %s", addr.String())
continue
}
if err = destinationConn.writePacket(buf[:n], udpAddr); err != nil {
if err = destinationConn.writePacket(buf[:n], addr); err != nil {
m.params.Logger.Errorf("could not write packet: %v", err)
}
}
}
func (m *UDPMuxDefault) getConn(ufrag string, isIPv6 bool) (val *udpMuxedConn, ok bool) {
func (m *UDPMuxDefault) getConn(ufrag string, isIPv6 bool, local net.IP) (val *udpMuxedConn, ok bool) {
var conns map[ipAddr]*udpMuxedConn
if isIPv6 {
val, ok = m.connsIPv6[ufrag]
conns, ok = m.connsIPv6[ufrag]
} else {
val, ok = m.connsIPv4[ufrag]
conns, ok = m.connsIPv4[ufrag]
}
if conns != nil {
val, ok = conns[ipAddr(local.String())]
}
return
}
+5 -5
View File
@@ -10,7 +10,7 @@ import "net"
// a UDPMux, in which case it will return a single connection for one
// of the ports.
type AllConnsGetter interface {
GetAllConns(ufrag string, isIPv6 bool) ([]net.PacketConn, error)
GetAllConns(ufrag string, isIPv6 bool, local net.IP) ([]net.PacketConn, error)
}
// MultiUDPMuxDefault implements both UDPMux and AllConnsGetter,
@@ -32,13 +32,13 @@ func NewMultiUDPMuxDefault(muxs ...UDPMux) *MultiUDPMuxDefault {
// creates the connection if an existing one can't be found. This, unlike
// GetAllConns, will only return a single PacketConn from the first
// mux that was passed in to NewMultiUDPMuxDefault.
func (m *MultiUDPMuxDefault) GetConn(ufrag string, isIPv6 bool) (net.PacketConn, error) {
func (m *MultiUDPMuxDefault) GetConn(ufrag string, isIPv6 bool, local net.IP) (net.PacketConn, error) {
// NOTE: We always use the first element here in order to maintain the
// behavior of using an existing connection if one exists.
if len(m.muxs) == 0 {
return nil, errNoUDPMuxAvailable
}
return m.muxs[0].GetConn(ufrag, isIPv6)
return m.muxs[0].GetConn(ufrag, isIPv6, local)
}
// RemoveConnByUfrag stops and removes the muxed packet connection
@@ -50,14 +50,14 @@ func (m *MultiUDPMuxDefault) RemoveConnByUfrag(ufrag string) {
}
// GetAllConns returns a PacketConn for each underlying UDPMux
func (m *MultiUDPMuxDefault) GetAllConns(ufrag string, isIPv6 bool) ([]net.PacketConn, error) {
func (m *MultiUDPMuxDefault) GetAllConns(ufrag string, isIPv6 bool, local net.IP) ([]net.PacketConn, error) {
if len(m.muxs) == 0 {
// Make sure that we either return at least one connection or an error.
return nil, errNoUDPMuxAvailable
}
var conns []net.PacketConn
for _, mux := range m.muxs {
conn, err := mux.GetConn(ufrag, isIPv6)
conn, err := mux.GetConn(ufrag, isIPv6, local)
if err != nil {
// For now, this implementation is all or none.
return nil, err
+10 -4
View File
@@ -60,12 +60,12 @@ func TestMultiUDPMux(t *testing.T) {
require.NoError(t, udpMuxMulti.Close())
// can't create more connections
_, err = udpMuxMulti.GetConn("failufrag", false)
_, err = udpMuxMulti.GetConn("failufrag", false, net.IP{})
require.Error(t, err)
}
func testMultiUDPMuxConnections(t *testing.T, udpMuxMulti *MultiUDPMuxDefault, ufrag string, network string) {
pktConns, err := udpMuxMulti.GetAllConns(ufrag, false)
pktConns, err := udpMuxMulti.GetAllConns(ufrag, false, net.IP{127, 0, 0, 1})
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
for _, c := range pktConns {
@@ -75,11 +75,17 @@ func testMultiUDPMuxConnections(t *testing.T, udpMuxMulti *MultiUDPMuxDefault, u
require.Len(t, pktConns, len(udpMuxMulti.muxs), "there should be a PacketConn for every mux")
// Try talking with each PacketConn
for _, pktConn := range pktConns {
for i, pktConn := range pktConns {
remoteConn, err := net.DialUDP(network, nil, &net.UDPAddr{
Port: pktConn.LocalAddr().(*net.UDPAddr).Port,
})
require.NoError(t, err, "error dialing test udp connection")
testMuxConnectionPair(t, pktConn, remoteConn, ufrag)
localConn, err := udpMuxMulti.muxs[i].GetConn(ufrag, false, remoteConn.RemoteAddr().(*net.UDPAddr).IP)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
}()
testMuxConnectionPair(t, localConn, remoteConn, ufrag)
}
}
+8 -8
View File
@@ -40,7 +40,7 @@ func TestUDPMux(t *testing.T) {
_ = conn.Close()
}()
require.NotNil(t, udpMux.LocalAddr(), "tcpMux.LocalAddr() is nil")
require.NotNil(t, udpMux.LocalAddr(), "udpMux.LocalAddr() is nil")
wg := sync.WaitGroup{}
@@ -66,7 +66,7 @@ func TestUDPMux(t *testing.T) {
require.NoError(t, udpMux.Close())
// can't create more connections
_, err = udpMux.GetConn("failufrag", false)
_, err = udpMux.GetConn("failufrag", false, net.IPv4zero)
require.Error(t, err)
}
@@ -111,17 +111,17 @@ func TestAddressEncoding(t *testing.T) {
}
func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string, network string) {
pktConn, err := udpMux.GetConn(ufrag, false)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
}()
remoteConn, err := net.DialUDP(network, nil, &net.UDPAddr{
Port: udpMux.LocalAddr().(*net.UDPAddr).Port,
})
require.NoError(t, err, "error dialing test udp connection")
pktConn, err := udpMux.GetConn(ufrag, false, remoteConn.RemoteAddr().(*net.UDPAddr).IP)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
}()
testMuxConnectionPair(t, pktConn, remoteConn, ufrag)
}
+23 -19
View File
@@ -33,7 +33,7 @@ type UniversalUDPMuxDefault struct {
// UniversalUDPMuxParams are parameters for UniversalUDPMux server reflexive.
type UniversalUDPMuxParams struct {
Logger logging.LeveledLogger
UDPConn net.PacketConn
UDPConn UDPMuxConn
XORMappedAddrCacheTTL time.Duration
}
@@ -54,7 +54,7 @@ func NewUniversalUDPMuxDefault(params UniversalUDPMuxParams) *UniversalUDPMuxDef
// wrap UDP connection, process server reflexive messages
// before they are passed to the UDPMux connection handler (connWorker)
m.params.UDPConn = &udpConn{
PacketConn: params.UDPConn,
UDPMuxConn: params.UDPConn,
mux: m,
logger: params.Logger,
}
@@ -71,7 +71,7 @@ func NewUniversalUDPMuxDefault(params UniversalUDPMuxParams) *UniversalUDPMuxDef
// udpConn is a wrapper around UDPMux conn that overrides ReadFrom and handles STUN/TURN packets
type udpConn struct {
net.PacketConn
UDPMuxConn
mux *UniversalUDPMuxDefault
logger logging.LeveledLogger
}
@@ -85,43 +85,47 @@ func (m *UniversalUDPMuxDefault) GetRelayedAddr(turnAddr net.Addr, deadline time
// GetConnForURL add uniques to the muxed connection by concatenating ufrag and URL (e.g. STUN URL) to be able to support multiple STUN/TURN servers
// and return a unique connection per server.
func (m *UniversalUDPMuxDefault) GetConnForURL(ufrag string, url string, isIPv6 bool) (net.PacketConn, error) {
return m.UDPMuxDefault.GetConn(fmt.Sprintf("%s%s", ufrag, url), isIPv6)
return m.UDPMuxDefault.GetConn(fmt.Sprintf("%s%s", ufrag, url), isIPv6, net.IPv4zero)
}
// ReadFrom is called by UDPMux connWorker and handles packets coming from the STUN server discovering a mapped address.
// ReadMsgUDP is called by UDPMux connWorker and handles packets coming from the STUN server discovering a mapped address.
// It passes processed packets further to the UDPMux (maybe this is not really necessary).
func (c *udpConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
n, addr, err = c.PacketConn.ReadFrom(p)
func (c *udpConn) ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net.UDPAddr, err error) {
n, oobn, flags, addr, err = c.UDPMuxConn.ReadMsgUDP(b, oob)
if err != nil {
return
}
if stun.IsMessage(p[:n]) {
if stun.IsMessage(b[:n]) {
bytes := make([]byte, n)
copy(bytes, b[:n])
msg := &stun.Message{
Raw: append([]byte{}, p[:n]...),
Raw: bytes,
}
if err = msg.Decode(); err != nil {
c.logger.Warnf("Failed to handle decode ICE from %s: %v", addr.String(), err)
return n, addr, nil
}
udpAddr, ok := addr.(*net.UDPAddr)
if !ok {
// message about this err will be logged in the UDPMux
err = nil
return
}
if c.mux.isXORMappedResponse(msg, udpAddr.String()) {
err = c.mux.handleXORMappedResponse(udpAddr, msg)
if c.mux.isXORMappedResponse(msg, addr.String()) {
err = c.mux.handleXORMappedResponse(addr, msg)
if err != nil {
c.logger.Debugf("%w: %v", errGetXorMappedAddrResponse, err)
return n, addr, nil
err = nil
return
}
return
}
}
return n, addr, err
return
}
func (c *udpConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
oob := make([]byte, 100)
n, _, _, addr, err = c.ReadMsgUDP(p, oob)
return
}
// isXORMappedResponse indicates whether the message is a XORMappedAddress and is coming from the known STUN server.
+1 -1
View File
@@ -41,7 +41,7 @@ func TestUniversalUDPMux(t *testing.T) {
}
func testMuxSrflxConnection(t *testing.T, udpMux *UniversalUDPMuxDefault, ufrag string, network string) {
pktConn, err := udpMux.GetConn(ufrag, false)
pktConn, err := udpMux.GetConn(ufrag, false, net.IPv4zero)
require.NoError(t, err, "error retrieving muxed connection for ufrag")
defer func() {
_ = pktConn.Close()
+1
View File
@@ -16,6 +16,7 @@ type udpMuxedConnParams struct {
AddrPool *sync.Pool
Key string
LocalAddr net.Addr
LocalIP net.IP
Logger logging.LeveledLogger
}
+51
View File
@@ -0,0 +1,51 @@
//go:build !js
package ice
import (
"bytes"
"encoding/binary"
"errors"
"net"
"syscall"
"github.com/pion/logging"
)
var errUnknownOobData = errors.New("unknown oob data")
func setUDPSocketOptionsForLocalAddr(fd uintptr, logger logging.LeveledLogger) {
if err := syscall.SetsockoptInt(int(fd), syscall.IPPROTO_IPV6, syscall.IPV6_2292PKTINFO, 1); err != nil {
logger.Warnf("Failed to set sockopt IPV6_2292PKTINFO: %s", err)
}
if err := syscall.SetsockoptInt(int(fd), syscall.IPPROTO_IP, syscall.IP_PKTINFO, 1); err != nil {
logger.Warnf("Failed to set sockopt IP_PKTINFO: %s", err)
}
}
func getLocalAddrFromOob(oob []byte) (net.IP, error) {
var localHost net.IP
// get destination local addr from received packet
oobBuffer := bytes.NewBuffer(oob)
msg := syscall.Cmsghdr{}
err := binary.Read(oobBuffer, binary.LittleEndian, &msg)
if err == nil {
switch {
case msg.Level == syscall.IPPROTO_IP && msg.Type == syscall.IP_PKTINFO:
packetInfo := syscall.Inet4Pktinfo{}
if err = binary.Read(oobBuffer, binary.LittleEndian, &packetInfo); err == nil {
localHost = net.IP(packetInfo.Addr[:])
return localHost, nil
}
case msg.Level == syscall.IPPROTO_IPV6 && msg.Type == syscall.IPV6_2292PKTINFO:
packetInfo := syscall.Inet6Pktinfo{}
if err = binary.Read(oobBuffer, binary.LittleEndian, &packetInfo); err == nil {
localHost = net.IP(packetInfo.Addr[:])
return localHost, nil
}
default:
return localHost, errUnknownOobData
}
}
return localHost, err
}
+19
View File
@@ -0,0 +1,19 @@
//go:build js
package ice
import (
"errors"
"net"
"github.com/pion/logging"
)
var errUnsupported = errors.New("unsupported")
func setUDPSocketOptionsForLocalAddr(fd uintptr, logger logging.LeveledLogger) {
}
func getLocalAddrFromOob(oob []byte) (net.IP, error) {
return nil, errUnsupported
}