diff --git a/agent_udpmux_test.go b/agent_udpmux_test.go index 11ec79e..f0930f0 100644 --- a/agent_udpmux_test.go +++ b/agent_udpmux_test.go @@ -23,66 +23,72 @@ func TestMuxAgent(t *testing.T) { const muxPort = 7686 - c, err := net.ListenUDP("udp4", &net.UDPAddr{ - IP: net.IPv4(127, 0, 0, 1), - Port: muxPort, - }) + caseAddrs := map[string]*net.UDPAddr{ + "unspecified": {Port: muxPort}, + "ipv4Loopback": {IP: net.IPv4(127, 0, 0, 1), Port: muxPort}, + } - require.NoError(t, err) + for subTest, addr := range caseAddrs { + muxAddr := addr + t.Run(subTest, func(t *testing.T) { + c, err := net.ListenUDP("udp", muxAddr) + require.NoError(t, err) - loggerFactory := logging.NewDefaultLoggerFactory() - udpMux, err := NewUDPMuxDefault(UDPMuxParams{ - Logger: loggerFactory.NewLogger("ice"), - UDPConn: c, - }) + loggerFactory := logging.NewDefaultLoggerFactory() + udpMux, err := NewUDPMuxDefault(UDPMuxParams{ + Logger: loggerFactory.NewLogger("ice"), + UDPConn: c, + }) - require.NoError(t, err) - defer func() { - _ = udpMux.Close() - _ = c.Close() - }() + require.NoError(t, err) + defer func() { + _ = udpMux.Close() + _ = c.Close() + }() - muxedA, err := NewAgent(&AgentConfig{ - UDPMux: udpMux, - CandidateTypes: []CandidateType{CandidateTypeHost}, - NetworkTypes: []NetworkType{ - NetworkTypeUDP4, - }, - }) - require.NoError(t, err) + muxedA, err := NewAgent(&AgentConfig{ + UDPMux: udpMux, + CandidateTypes: []CandidateType{CandidateTypeHost}, + NetworkTypes: []NetworkType{ + NetworkTypeUDP4, + }, + }) + require.NoError(t, err) - a, err := NewAgent(&AgentConfig{ - CandidateTypes: []CandidateType{CandidateTypeHost}, - NetworkTypes: supportedNetworkTypes(), - }) - require.NoError(t, err) + a, err := NewAgent(&AgentConfig{ + CandidateTypes: []CandidateType{CandidateTypeHost}, + NetworkTypes: supportedNetworkTypes(), + }) + require.NoError(t, err) - conn, muxedConn := connect(a, muxedA) + conn, muxedConn := connect(a, muxedA) - pair := muxedA.getSelectedPair() - require.NotNil(t, pair) - require.Equal(t, muxPort, pair.Local.Port()) + pair := muxedA.getSelectedPair() + require.NotNil(t, pair) + require.Equal(t, muxPort, pair.Local.Port()) - // send a packet to Mux - data := []byte("hello world") - _, err = conn.Write(data) - require.NoError(t, err) + // send a packet to Mux + data := []byte("hello world") + _, err = conn.Write(data) + require.NoError(t, err) - buffer := make([]byte, 1024) - n, err := muxedConn.Read(buffer) - require.NoError(t, err) - require.Equal(t, data, buffer[:n]) + buffer := make([]byte, 1024) + n, err := muxedConn.Read(buffer) + require.NoError(t, err) + require.Equal(t, data, buffer[:n]) - // send a packet from Mux - _, err = muxedConn.Write(data) - require.NoError(t, err) + // send a packet from Mux + _, err = muxedConn.Write(data) + require.NoError(t, err) - n, err = conn.Read(buffer) - require.NoError(t, err) - require.Equal(t, data, buffer[:n]) + n, err = conn.Read(buffer) + require.NoError(t, err) + require.Equal(t, data, buffer[:n]) - // close it down - require.NoError(t, conn.Close()) - require.NoError(t, muxedConn.Close()) - require.NoError(t, udpMux.Close()) + // close it down + require.NoError(t, conn.Close()) + require.NoError(t, muxedConn.Close()) + require.NoError(t, udpMux.Close()) + }) + } } diff --git a/udp_mux.go b/udp_mux.go index 85ac011..9914346 100644 --- a/udp_mux.go +++ b/udp_mux.go @@ -10,6 +10,7 @@ import ( "github.com/pion/logging" "github.com/pion/stun" + "github.com/pion/transport/vnet" ) // UDPMux allows multiple connections to go over a single UDP port @@ -37,6 +38,9 @@ type UDPMuxDefault struct { pool *sync.Pool mu sync.Mutex + + // for UDP connection listen at unspecified address + localAddrsForUnspecified []net.Addr } const maxAddrSize = 512 @@ -53,10 +57,34 @@ func NewUDPMuxDefault(params UDPMuxParams) (*UDPMuxDefault, error) { params.Logger = logging.NewDefaultLoggerFactory().NewLogger("ice") } + var localAddrsForUnspecified []net.Addr if addr, ok := params.UDPConn.LocalAddr().(*net.UDPAddr); !ok { return nil, errInvalidAddress } else if ok && addr.IP.IsUnspecified() { - return nil, errListenUnspecified + // return nil, errListenUnspecified + // For unspecified addresses, the correct behavior is to return errListenUnspecified, but + // it will break the applications that are already using unspecified UDP connection + // with UDPMuxDefault, so print a warn log and create a local address list for mux. + params.Logger.Warn("UDPMuxDefault should not listening on unspecified address, use NewMultiUDPMuxFromPort instead") + var networks []NetworkType + switch { + case addr.IP.To4() != nil: + networks = []NetworkType{NetworkTypeUDP4} + + case addr.IP.To16() != nil: + networks = []NetworkType{NetworkTypeUDP4, NetworkTypeUDP6} + + default: + return nil, errInvalidAddress + } + muxNet := vnet.NewNet(nil) + ips, err := localInterfaces(muxNet, nil, nil, networks) + if err != nil { + return nil, err + } + for _, ip := range ips { + localAddrsForUnspecified = append(localAddrsForUnspecified, &net.UDPAddr{IP: ip, Port: addr.Port}) + } } m := &UDPMuxDefault{ @@ -70,6 +98,7 @@ func NewUDPMuxDefault(params UDPMuxParams) (*UDPMuxDefault, error) { return newBufferHolder(receiveMTU + maxAddrSize) }, }, + localAddrsForUnspecified: localAddrsForUnspecified, } go m.connWorker() @@ -84,13 +113,18 @@ func (m *UDPMuxDefault) LocalAddr() net.Addr { // GetListenAddresses returns the list of addresses that this mux is listening on func (m *UDPMuxDefault) GetListenAddresses() []net.Addr { + if len(m.localAddrsForUnspecified) > 0 { + return m.localAddrsForUnspecified + } + return []net.Addr{m.LocalAddr()} } // 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, addr net.Addr) (net.PacketConn, error) { - if m.params.UDPConn.LocalAddr() != addr { + // don't check addr for mux using unspecified address + if len(m.localAddrsForUnspecified) == 0 && m.params.UDPConn.LocalAddr() != addr { return nil, errInvalidAddress } m.mu.Lock() diff --git a/udp_mux_test.go b/udp_mux_test.go index 557e453..16d44d1 100644 --- a/udp_mux_test.go +++ b/udp_mux_test.go @@ -78,18 +78,67 @@ func TestUDPMux(t *testing.T) { } } -func TestCantMuxUnspecifiedAddr(t *testing.T) { - conn, err := net.ListenUDP(udp, &net.UDPAddr{}) +func TestUDPMuxUnspecifiedAddr(t *testing.T) { + report := test.CheckRoutines(t) + defer report() + + lim := test.TimeOut(time.Second * 30) + defer lim.Stop() + + conn, err := net.ListenUDP(udp, nil) require.NoError(t, err) - _, err = NewUDPMuxDefault(UDPMuxParams{ - Logger: nil, - UDPConn: conn, - }) + conn4, err := net.ListenUDP(udp, &net.UDPAddr{IP: net.IPv4zero}) + require.NoError(t, err) - require.Equal(t, errListenUnspecified, err) + conn6, err := net.ListenUDP(udp, &net.UDPAddr{IP: net.IPv6unspecified}) + if err != nil { + t.Log("IPv6 is not supported on this machine") + } - _ = conn.Close() + for network, c := range map[string]net.PacketConn{udp: conn, udp4: conn4, udp6: conn6} { + if udpConn, ok := c.(*net.UDPConn); !ok || udpConn == nil { + continue + } + conn := c + t.Run(network, func(t *testing.T) { + udpMux, err := NewUDPMuxDefault(UDPMuxParams{ + Logger: nil, + UDPConn: conn, + }) + + require.NoError(t, err) + + defer func() { + _ = udpMux.Close() + _ = conn.Close() + }() + + require.NotNil(t, udpMux.LocalAddr(), "udpMux.LocalAddr() is nil") + + wg := sync.WaitGroup{} + + wg.Add(1) + go func() { + defer wg.Done() + testMuxConnection(t, udpMux, "ufrag1", udp) + }() + + // skip ipv6 test on i386 + const ptrSize = 32 << (^uintptr(0) >> 63) + if ptrSize != 32 || network != udp6 { + testMuxConnection(t, udpMux, "ufrag2", network) + } + + wg.Wait() + + require.NoError(t, udpMux.Close()) + + // can't create more connections + _, err = udpMux.GetConn("failufrag", udpMux.LocalAddr()) + require.Error(t, err) + }) + } } func TestAddressEncoding(t *testing.T) { @@ -139,7 +188,12 @@ func testMuxConnection(t *testing.T, udpMux *UDPMuxDefault, ufrag string, networ _ = pktConn.Close() }() - remoteConn, err := net.DialUDP(network, nil, pktConn.LocalAddr().(*net.UDPAddr)) + addr, ok := pktConn.LocalAddr().(*net.UDPAddr) + require.True(t, ok, "pktConn.LocalAddr() is not a net.UDPAddr") + if addr.IP.IsUnspecified() { + addr = &net.UDPAddr{Port: addr.Port} + } + remoteConn, err := net.DialUDP(network, nil, addr) require.NoError(t, err, "error dialing test udp connection") testMuxConnectionPair(t, pktConn, remoteConn, ufrag)