diff --git a/gather.go b/gather.go index 7dc8938..83cc8d4 100644 --- a/gather.go +++ b/gather.go @@ -110,21 +110,24 @@ func (a *Agent) gatherCandidates() <-chan struct{} { } <-gatherStateUpdated + var wg sync.WaitGroup for _, t := range a.candidateTypes { switch t { case CandidateTypeHost: a.gatherCandidatesLocal(a.networkTypes) case CandidateTypeServerReflexive: - a.gatherCandidatesSrflx(a.urls, a.networkTypes) + a.gatherCandidatesSrflx(a.urls, a.networkTypes, &wg) if a.extIPMapper != nil && a.extIPMapper.candidateType == CandidateTypeServerReflexive { - a.gatherCandidatesSrflxMapped(a.networkTypes) + a.gatherCandidatesSrflxMapped(a.networkTypes, &wg) } case CandidateTypeRelay: - if err := a.gatherCandidatesRelay(a.urls); err != nil { + if err := a.gatherCandidatesRelay(a.urls, &wg); err != nil { a.log.Errorf("Failed to gather relay candidates: %v\n", err) } } } + // Block until all STUN and TURN URLs have been gathered (or timed out) + wg.Wait() if err := a.run(func(agent *Agent) { closeChanCandidateOnce.Do(func() { close(agent.chanCandidate) @@ -199,52 +202,54 @@ func (a *Agent) gatherCandidatesLocal(networkTypes []NetworkType) { } } -func (a *Agent) gatherCandidatesSrflxMapped(networkTypes []NetworkType) { +func (a *Agent) gatherCandidatesSrflxMapped(networkTypes []NetworkType, wg *sync.WaitGroup) { for _, networkType := range networkTypes { network := networkType.String() - - 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) - continue - } - - laddr := conn.LocalAddr().(*net.UDPAddr) - mappedIP, err := a.extIPMapper.findExternalIP(laddr.IP.String()) - if err != nil { - closeConnAndLog(conn, a.log, fmt.Sprintf("1:1 NAT mapping is enabled but no external IP is found for %s\n", laddr.IP.String())) - continue - } - - srflxConfig := CandidateServerReflexiveConfig{ - Network: network, - Address: mappedIP.String(), - Port: laddr.Port, - Component: ComponentRTP, - RelAddr: laddr.IP.String(), - RelPort: laddr.Port, - } - c, err := NewCandidateServerReflexive(&srflxConfig) - if err != nil { - closeConnAndLog(conn, a.log, fmt.Sprintf("Failed to create server reflexive candidate: %s %s %d: %v\n", - network, - mappedIP.String(), - laddr.Port, - err)) - continue - } - - if err := a.addCandidate(c, conn); err != nil { - if closeErr := c.close(); closeErr != nil { - a.log.Warnf("Failed to close candidate: %v", closeErr) + 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) + return } - a.log.Warnf("Failed to append to localCandidates and run onCandidateHdlr: %v\n", err) - } + + laddr := conn.LocalAddr().(*net.UDPAddr) + mappedIP, err := a.extIPMapper.findExternalIP(laddr.IP.String()) + if err != nil { + closeConnAndLog(conn, a.log, fmt.Sprintf("1:1 NAT mapping is enabled but no external IP is found for %s\n", laddr.IP.String())) + return + } + + srflxConfig := CandidateServerReflexiveConfig{ + Network: network, + Address: mappedIP.String(), + Port: laddr.Port, + Component: ComponentRTP, + RelAddr: laddr.IP.String(), + RelPort: laddr.Port, + } + c, err := NewCandidateServerReflexive(&srflxConfig) + if err != nil { + closeConnAndLog(conn, a.log, fmt.Sprintf("Failed to create server reflexive candidate: %s %s %d: %v\n", + network, + mappedIP.String(), + laddr.Port, + err)) + return + } + + if err := a.addCandidate(c, conn); err != nil { + if closeErr := c.close(); closeErr != nil { + a.log.Warnf("Failed to close candidate: %v", closeErr) + } + a.log.Warnf("Failed to append to localCandidates and run onCandidateHdlr: %v\n", err) + } + }() } } -func (a *Agent) gatherCandidatesSrflx(urls []*URL, networkTypes []NetworkType) { - var wg sync.WaitGroup +func (a *Agent) gatherCandidatesSrflx(urls []*URL, networkTypes []NetworkType, wg *sync.WaitGroup) { for _, networkType := range networkTypes { for i := range urls { if urls[i].Scheme != SchemeTypeSTUN { @@ -300,14 +305,9 @@ func (a *Agent) gatherCandidatesSrflx(urls []*URL, networkTypes []NetworkType) { }(*urls[i], networkType.String()) } } - - // Block until all STUN URLs have been gathered (or timed out) - wg.Wait() } -func (a *Agent) gatherCandidatesRelay(urls []*URL) error { - var wg sync.WaitGroup - +func (a *Agent) gatherCandidatesRelay(urls []*URL, wg *sync.WaitGroup) error { network := NetworkTypeUDP4.String() // TODO IPv6 for i := range urls { switch { @@ -448,7 +448,5 @@ func (a *Agent) gatherCandidatesRelay(urls []*URL) error { }(*urls[i]) } - // Block until all STUN URLs have been gathered (or timed out) - wg.Wait() return nil } diff --git a/gather_test.go b/gather_test.go index ffc84cf..b56927d 100644 --- a/gather_test.go +++ b/gather_test.go @@ -255,6 +255,73 @@ func TestTURNConcurrency(t *testing.T) { }) } +// Assert that STUN and TURN gathering are done concurrently +func TestSTUNTURNConcurrency(t *testing.T) { + report := test.CheckRoutines(t) + defer report() + + lim := test.TimeOut(time.Second * 8) + defer lim.Stop() + + serverPort := randomPort(t) + serverListener, err := net.ListenPacket("udp4", "127.0.0.1:"+strconv.Itoa(serverPort)) + assert.NoError(t, err) + + server, err := turn.NewServer(turn.ServerConfig{ + Realm: "pion.ly", + AuthHandler: optimisticAuthHandler, + PacketConnConfigs: []turn.PacketConnConfig{ + { + PacketConn: serverListener, + RelayAddressGenerator: &turn.RelayAddressGeneratorNone{Address: "127.0.0.1"}, + }, + }, + }) + assert.NoError(t, err) + + urls := []*URL{} + for i := 0; i <= 10; i++ { + urls = append(urls, &URL{ + Scheme: SchemeTypeSTUN, + Host: "127.0.0.1", + Port: serverPort + 1, + }) + } + urls = append(urls, &URL{ + Scheme: SchemeTypeTURN, + Proto: ProtoTypeUDP, + Host: "127.0.0.1", + Port: serverPort, + Username: "username", + Password: "password", + }) + + a, err := NewAgent(&AgentConfig{ + NetworkTypes: supportedNetworkTypes, + Urls: urls, + CandidateTypes: []CandidateType{CandidateTypeServerReflexive, CandidateTypeRelay}, + }) + assert.NoError(t, err) + + { + gatherLim := test.TimeOut(time.Second * 3) // As TURN and STUN should be checked in parallel, this should complete before the default STUN timeout (5s) + candidateGathered, candidateGatheredFunc := context.WithCancel(context.Background()) + assert.NoError(t, a.OnCandidate(func(c Candidate) { + if c != nil { + candidateGatheredFunc() + } + })) + assert.NoError(t, a.GatherCandidates()) + + <-candidateGathered.Done() + + gatherLim.Stop() + } + + assert.NoError(t, a.Close()) + assert.NoError(t, server.Close()) +} + func TestCloseConnLog(t *testing.T) { a, err := NewAgent(&AgentConfig{}) assert.NoError(t, err)