From e90a58e51aa8bc73e053cf70734e54c68813b06c Mon Sep 17 00:00:00 2001 From: cnderrauber Date: Tue, 22 Nov 2022 11:15:14 +0800 Subject: [PATCH] Add option to include loopback candidate Add option to include loopback candidate --- agent.go | 3 ++ agent_config.go | 5 ++- gather.go | 2 +- gather_test.go | 85 ++++++++++++++++++++++++++++++++++++++++++++- gather_vnet_test.go | 12 +++---- udp_mux.go | 2 +- udp_mux_multi.go | 12 ++++++- util.go | 6 ++-- 8 files changed, 113 insertions(+), 14 deletions(-) diff --git a/agent.go b/agent.go index c1af0d7..dcca6c5 100644 --- a/agent.go +++ b/agent.go @@ -130,6 +130,7 @@ type Agent struct { interfaceFilter func(string) bool ipFilter func(net.IP) bool + includeLoopback bool insecureSkipVerify bool @@ -317,6 +318,8 @@ func NewAgent(config *AgentConfig) (*Agent, error) { //nolint:gocognit ipFilter: config.IPFilter, insecureSkipVerify: config.InsecureSkipVerify, + + includeLoopback: config.IncludeLoopback, } a.tcpMux = config.TCPMux diff --git a/agent_config.go b/agent_config.go index 98bbbff..54a61ba 100644 --- a/agent_config.go +++ b/agent_config.go @@ -165,8 +165,11 @@ type AgentConfig struct { // dial interface in order to support corporate proxies ProxyDialer proxy.Dialer - // Accept aggressive nomination in RFC 5245 for compatible with chrome and other browsers + // Deprecated: AcceptAggressiveNomination always enabled. AcceptAggressiveNomination bool + + // Include loopback addresses in the candidate list. + IncludeLoopback bool } // initWithDefaults populates an agent and falls back to defaults if fields are unset diff --git a/gather.go b/gather.go index 6cde60a..c2f67a0 100644 --- a/gather.go +++ b/gather.go @@ -149,7 +149,7 @@ func (a *Agent) gatherCandidatesLocal(ctx context.Context, networkTypes []Networ delete(networks, udp) } - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, networkTypes) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, networkTypes, a.includeLoopback) if err != nil { a.log.Warnf("failed to iterate local interfaces, host candidates will not be gathered %s", err) return diff --git a/gather_test.go b/gather_test.go index ac595bd..55fdbc2 100644 --- a/gather_test.go +++ b/gather_test.go @@ -13,6 +13,7 @@ import ( "sort" "strconv" "sync" + "sync/atomic" "testing" "time" @@ -31,7 +32,7 @@ func TestListenUDP(t *testing.T) { a, err := NewAgent(&AgentConfig{}) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) assert.NotEqual(t, len(localIPs), 0, "localInterfaces found no interfaces, unable to test") assert.NoError(t, err) @@ -86,6 +87,88 @@ func TestListenUDP(t *testing.T) { assert.NoError(t, a.Close()) } +func TestLoopbackCandidate(t *testing.T) { + report := test.CheckRoutines(t) + defer report() + + lim := test.TimeOut(time.Second * 30) + defer lim.Stop() + type testCase struct { + name string + agentConfig *AgentConfig + loExpected bool + } + mux, err := NewMultiUDPMuxFromPort(12500) + assert.NoError(t, err) + muxWithLo, errlo := NewMultiUDPMuxFromPort(12501, UDPMuxFromPortWithLoopback()) + assert.NoError(t, errlo) + testCases := []testCase{ + { + name: "mux should not have loopback candidate", + agentConfig: &AgentConfig{ + NetworkTypes: []NetworkType{NetworkTypeUDP4, NetworkTypeUDP6}, + UDPMux: mux, + }, + loExpected: false, + }, + { + name: "mux with loopback should not have loopback candidate", + agentConfig: &AgentConfig{ + NetworkTypes: []NetworkType{NetworkTypeUDP4, NetworkTypeUDP6}, + UDPMux: muxWithLo, + }, + loExpected: true, + }, + { + name: "includeloopback enabled", + agentConfig: &AgentConfig{ + NetworkTypes: []NetworkType{NetworkTypeUDP4, NetworkTypeUDP6}, + IncludeLoopback: true, + }, + loExpected: true, + }, + { + name: "includeloopback disabled", + agentConfig: &AgentConfig{ + NetworkTypes: []NetworkType{NetworkTypeUDP4, NetworkTypeUDP6}, + IncludeLoopback: false, + }, + loExpected: false, + }, + } + + for _, tc := range testCases { + tcase := tc + t.Run(tcase.name, func(t *testing.T) { + a, err := NewAgent(tc.agentConfig) + assert.NoError(t, err) + + candidateGathered, candidateGatheredFunc := context.WithCancel(context.Background()) + var loopback int32 + assert.NoError(t, a.OnCandidate(func(c Candidate) { + if c != nil { + if net.ParseIP(c.Address()).IsLoopback() { + atomic.StoreInt32(&loopback, 1) + } + } else { + candidateGatheredFunc() + return + } + t.Log(c.NetworkType(), c.Priority(), c) + })) + assert.NoError(t, a.GatherCandidates()) + + <-candidateGathered.Done() + + assert.NoError(t, a.Close()) + assert.Equal(t, tcase.loExpected, atomic.LoadInt32(&loopback) == 1) + }) + } + + assert.NoError(t, mux.Close()) + assert.NoError(t, muxWithLo.Close()) +} + // Assert that STUN gathering is done concurrently func TestSTUNConcurrency(t *testing.T) { report := test.CheckRoutines(t) diff --git a/gather_vnet_test.go b/gather_vnet_test.go index 87d8130..49ec850 100644 --- a/gather_vnet_test.go +++ b/gather_vnet_test.go @@ -29,7 +29,7 @@ func TestVNetGather(t *testing.T) { }) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if len(localIPs) > 0 { t.Fatal("should return no local IP") } else if err != nil { @@ -69,7 +69,7 @@ func TestVNetGather(t *testing.T) { }) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if len(localIPs) == 0 { t.Fatal("should have one local IP") } else if err != nil { @@ -112,7 +112,7 @@ func TestVNetGather(t *testing.T) { t.Fatalf("Failed to create agent: %s", err) } - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if len(localIPs) == 0 { t.Fatal("localInterfaces found no interfaces, unable to test") } else if err != nil { @@ -385,7 +385,7 @@ func TestVNetGatherWithInterfaceFilter(t *testing.T) { }) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if err != nil { t.Fatal(err) } else if len(localIPs) != 0 { @@ -405,7 +405,7 @@ func TestVNetGatherWithInterfaceFilter(t *testing.T) { }) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if err != nil { t.Fatal(err) } else if len(localIPs) != 0 { @@ -425,7 +425,7 @@ func TestVNetGatherWithInterfaceFilter(t *testing.T) { }) assert.NoError(t, err) - localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}) + localIPs, err := localInterfaces(a.net, a.interfaceFilter, a.ipFilter, []NetworkType{NetworkTypeUDP4}, false) if err != nil { t.Fatal(err) } else if len(localIPs) == 0 { diff --git a/udp_mux.go b/udp_mux.go index 9cbd285..d589020 100644 --- a/udp_mux.go +++ b/udp_mux.go @@ -78,7 +78,7 @@ func NewUDPMuxDefault(params UDPMuxParams) *UDPMuxDefault { } if len(networks) > 0 { muxNet := vnet.NewNet(nil) - ips, err := localInterfaces(muxNet, nil, nil, networks) + ips, err := localInterfaces(muxNet, nil, nil, networks, true) if err == nil { for _, ip := range ips { localAddrsForUnspecified = append(localAddrsForUnspecified, &net.UDPAddr{IP: ip, Port: addr.Port}) diff --git a/udp_mux_multi.go b/udp_mux_multi.go index 47ce404..0c9b15b 100644 --- a/udp_mux_multi.go +++ b/udp_mux_multi.go @@ -81,7 +81,7 @@ func NewMultiUDPMuxFromPort(port int, opts ...UDPMuxFromPortOption) (*MultiUDPMu opt.apply(¶ms) } muxNet := vnet.NewNet(nil) - ips, err := localInterfaces(muxNet, params.ifFilter, params.ipFilter, params.networks) + ips, err := localInterfaces(muxNet, params.ifFilter, params.ipFilter, params.networks, params.includeLoopback) if err != nil { return nil, err } @@ -130,6 +130,7 @@ type multiUDPMuxFromPortParam struct { readBufferSize int writeBufferSize int logger logging.LeveledLogger + includeLoopback bool } type udpMuxFromPortOption struct { @@ -193,3 +194,12 @@ func UDPMuxFromPortWithLogger(logger logging.LeveledLogger) UDPMuxFromPortOption }, } } + +// UDPMuxFromPortWithLoopback set loopback interface should be included +func UDPMuxFromPortWithLoopback() UDPMuxFromPortOption { + return &udpMuxFromPortOption{ + f: func(p *multiUDPMuxFromPortParam) { + p.includeLoopback = true + }, + } +} diff --git a/util.go b/util.go index 22260de..7321a2e 100644 --- a/util.go +++ b/util.go @@ -132,7 +132,7 @@ func stunRequest(read func([]byte) (int, error), write func([]byte) (int, error) return res, nil } -func localInterfaces(vnet *vnet.Net, interfaceFilter func(string) bool, ipFilter func(net.IP) bool, networkTypes []NetworkType) ([]net.IP, error) { //nolint:gocognit +func localInterfaces(vnet *vnet.Net, interfaceFilter func(string) bool, ipFilter func(net.IP) bool, networkTypes []NetworkType, includeLoopback bool) ([]net.IP, error) { //nolint:gocognit ips := []net.IP{} ifaces, err := vnet.Interfaces() if err != nil { @@ -154,7 +154,7 @@ func localInterfaces(vnet *vnet.Net, interfaceFilter func(string) bool, ipFilter if iface.Flags&net.FlagUp == 0 { continue // interface down } - if iface.Flags&net.FlagLoopback != 0 { + if (iface.Flags&net.FlagLoopback != 0) && !includeLoopback { continue // loopback interface } @@ -175,7 +175,7 @@ func localInterfaces(vnet *vnet.Net, interfaceFilter func(string) bool, ipFilter case *net.IPAddr: ip = addr.IP } - if ip == nil || ip.IsLoopback() { + if ip == nil || (ip.IsLoopback() && !includeLoopback) { continue }