Gather TURN/STUN concurrently

Each URL is gathered in a goroutine. Add tests to ensure that this
doesn't regress in the future.

Resolves #118
This commit is contained in:
Sean DuBois
2020-02-23 22:05:13 -08:00
parent 50bd9f60bd
commit 3e38db1ea5
3 changed files with 316 additions and 160 deletions
+3 -9
View File
@@ -73,13 +73,7 @@ func TestServerReflexiveOnlyConnection(t *testing.T) {
<-aConnected
<-bConnected
if err = aAgent.Close(); err != nil {
t.Fatal(err)
}
if err = bAgent.Close(); err != nil {
t.Fatal(err)
}
if err = server.Close(); err != nil {
t.Fatal(err)
}
assert.NoError(t, aAgent.Close())
assert.NoError(t, bAgent.Close())
assert.NoError(t, server.Close())
}
+157 -151
View File
@@ -3,6 +3,7 @@ package ice
import (
"fmt"
"net"
"sync"
"time"
"github.com/pion/logging"
@@ -39,7 +40,6 @@ func (a *Agent) GatherCandidates() error {
}
go a.gatherCandidates()
gatherErrChan <- nil
})
if runErr != nil {
@@ -197,178 +197,184 @@ func (a *Agent) gatherCandidatesSrflxMapped(networkTypes []NetworkType) {
}
func (a *Agent) gatherCandidatesSrflx(urls []*URL, networkTypes []NetworkType) {
var stunURLs []*URL
for _, url := range urls {
if url.Scheme == SchemeTypeSTUN {
stunURLs = append(stunURLs, url)
var wg sync.WaitGroup
for _, networkType := range networkTypes {
for i := range urls {
if urls[i].Scheme != SchemeTypeSTUN {
continue
}
wg.Add(1)
go func(url URL, network string) {
defer wg.Done()
hostPort := fmt.Sprintf("%s:%d", url.Host, url.Port)
serverAddr, err := a.net.ResolveUDPAddr(network, hostPort)
if err != nil {
a.log.Warnf("failed to resolve stun host: %s: %v", hostPort, err)
return
}
conn, err := listenUDPInPortRange(a.net, a.log, int(a.portmax), int(a.portmin), network, &net.UDPAddr{IP: nil, Port: 0})
if err != nil {
closeConnAndLog(conn, a.log, fmt.Sprintf("Failed to listen for %s: %v\n", serverAddr.String(), err))
return
}
xoraddr, err := getXORMappedAddr(conn, serverAddr, stunGatherTimeout)
if err != nil {
closeConnAndLog(conn, a.log, fmt.Sprintf("could not get server reflexive address %s %s: %v\n", network, url, err))
return
}
ip := xoraddr.IP
port := xoraddr.Port
laddr := conn.LocalAddr().(*net.UDPAddr)
srflxConfig := CandidateServerReflexiveConfig{
Network: network,
Address: ip.String(),
Port: 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, ip, 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)
}
}(*urls[i], networkType.String())
}
}
for _, networkType := range networkTypes {
network := networkType.String()
for _, url := range stunURLs {
if url.Scheme != SchemeTypeSTUN {
continue
// Block until all STUN URLs have been gathered (or timed out)
wg.Wait()
}
func (a *Agent) gatherCandidatesRelay(urls []*URL) error {
var wg sync.WaitGroup
network := NetworkTypeUDP4.String() // TODO IPv6
for i := range urls {
switch {
case urls[i].Scheme != SchemeTypeTURN:
continue
case urls[i].Username == "":
return ErrUsernameEmpty
case urls[i].Password == "":
return ErrPasswordEmpty
}
wg.Add(1)
go func(url URL) {
defer wg.Done()
TURNServerAddr := fmt.Sprintf("%s:%d", url.Host, url.Port)
var (
locConn net.PacketConn
err error
RelAddr string
RelPort int
)
if url.Proto == ProtoTypeUDP {
locConn, err = a.net.ListenPacket(network, "0.0.0.0:0")
if err != nil {
a.log.Warnf("Failed to listen %s: %v\n", network, err)
return
}
RelAddr = locConn.LocalAddr().(*net.UDPAddr).IP.String()
RelPort = locConn.LocalAddr().(*net.UDPAddr).Port
} else {
var (
tcpAddr *net.TCPAddr
tcpConn *net.TCPConn
)
tcpAddr, err = net.ResolveTCPAddr(NetworkTypeTCP4.String(), TURNServerAddr)
if err != nil {
a.log.Warnf("Failed to resolve TCP Addr %s: %v\n", TURNServerAddr, err)
return
}
tcpConn, err = net.DialTCP(NetworkTypeTCP4.String(), nil, tcpAddr)
if err != nil {
a.log.Warnf("Failed to Dial TCP Addr %s: %v\n", TURNServerAddr, err)
return
}
RelAddr = tcpConn.LocalAddr().(*net.TCPAddr).IP.String()
RelPort = tcpConn.LocalAddr().(*net.TCPAddr).Port
locConn = turn.NewSTUNConn(tcpConn)
}
hostPort := fmt.Sprintf("%s:%d", url.Host, url.Port)
serverAddr, err := a.net.ResolveUDPAddr(network, hostPort)
client, err := turn.NewClient(&turn.ClientConfig{
TURNServerAddr: TURNServerAddr,
Conn: locConn,
Username: url.Username,
Password: url.Password,
LoggerFactory: a.loggerFactory,
Net: a.net,
})
if err != nil {
a.log.Warnf("failed to resolve stun host: %s: %v", hostPort, err)
continue
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to build new turn.Client %s %s\n", TURNServerAddr, err))
return
}
conn, err := listenUDPInPortRange(a.net, a.log, int(a.portmax), int(a.portmin), network, &net.UDPAddr{IP: nil, Port: 0})
if err = client.Listen(); err != nil {
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to listen on turn.Client %s %s\n", TURNServerAddr, err))
return
}
relayConn, err := client.Allocate()
if err != nil {
closeConnAndLog(conn, a.log, fmt.Sprintf("Failed to listen for %s: %v\n", serverAddr.String(), err))
continue
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to allocate on turn.Client %s %s\n", TURNServerAddr, err))
return
}
xoraddr, err := getXORMappedAddr(conn, serverAddr, stunGatherTimeout)
if err != nil {
closeConnAndLog(conn, a.log, fmt.Sprintf("could not get server reflexive address %s %s: %v\n", network, url, err))
continue
}
ip := xoraddr.IP
port := xoraddr.Port
laddr := conn.LocalAddr().(*net.UDPAddr)
srflxConfig := CandidateServerReflexiveConfig{
raddr := relayConn.LocalAddr().(*net.UDPAddr)
relayConfig := CandidateRelayConfig{
Network: network,
Address: ip.String(),
Port: port,
Component: ComponentRTP,
RelAddr: laddr.IP.String(),
RelPort: laddr.Port,
Address: raddr.IP.String(),
Port: raddr.Port,
RelAddr: RelAddr,
RelPort: RelPort,
OnClose: func() error {
client.Close()
return locConn.Close()
},
}
c, err := NewCandidateServerReflexive(&srflxConfig)
candidate, err := NewCandidateRelay(&relayConfig)
if err != nil {
closeConnAndLog(conn, a.log, fmt.Sprintf("Failed to create server reflexive candidate: %s %s %d: %v\n", network, ip, port, err))
continue
if relayConErr := relayConn.Close(); relayConErr != nil {
a.log.Warnf("Failed to close relay %v", relayConErr)
}
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to create relay candidate: %s %s: %v\n", network, raddr.String(), err))
return
}
if err := a.addCandidate(c, conn); err != nil {
if closeErr := c.close(); closeErr != nil {
if err := a.addCandidate(candidate, relayConn); err != nil {
if closeErr := candidate.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) gatherCandidatesRelay(urls []*URL) error {
for _, url := range urls {
switch {
case url.Scheme != SchemeTypeTURN:
continue
case url.Username == "":
return ErrUsernameEmpty
case url.Password == "":
return ErrPasswordEmpty
}
}
network := NetworkTypeUDP4.String() // TODO IPv6
for _, url := range urls {
TURNServerAddr := fmt.Sprintf("%s:%d", url.Host, url.Port)
var (
locConn net.PacketConn
err error
RelAddr string
RelPort int
)
if url.Proto == ProtoTypeUDP {
locConn, err = a.net.ListenPacket(network, "0.0.0.0:0")
if err != nil {
a.log.Warnf("Failed to listen %s: %v\n", network, err)
continue
}
RelAddr = locConn.LocalAddr().(*net.UDPAddr).IP.String()
RelPort = locConn.LocalAddr().(*net.UDPAddr).Port
} else {
var (
tcpAddr *net.TCPAddr
tcpConn *net.TCPConn
)
tcpAddr, err = net.ResolveTCPAddr(NetworkTypeTCP4.String(), TURNServerAddr)
if err != nil {
a.log.Warnf("Failed to resolve TCP Addr %s: %v\n", TURNServerAddr, err)
continue
}
tcpConn, err = net.DialTCP(NetworkTypeTCP4.String(), nil, tcpAddr)
if err != nil {
a.log.Warnf("Failed to Dial TCP Addr %s: %v\n", TURNServerAddr, err)
continue
}
RelAddr = tcpConn.LocalAddr().(*net.TCPAddr).IP.String()
RelPort = tcpConn.LocalAddr().(*net.TCPAddr).Port
locConn = turn.NewSTUNConn(tcpConn)
}
client, err := turn.NewClient(&turn.ClientConfig{
TURNServerAddr: TURNServerAddr,
Conn: locConn,
Username: url.Username,
Password: url.Password,
LoggerFactory: a.loggerFactory,
Net: a.net,
})
if err != nil {
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to build new turn.Client %s %s\n", TURNServerAddr, err))
continue
}
if err = client.Listen(); err != nil {
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to listen on turn.Client %s %s\n", TURNServerAddr, err))
continue
}
relayConn, err := client.Allocate()
if err != nil {
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to allocate on turn.Client %s %s\n", TURNServerAddr, err))
continue
}
raddr := relayConn.LocalAddr().(*net.UDPAddr)
relayConfig := CandidateRelayConfig{
Network: network,
Component: ComponentRTP,
Address: raddr.IP.String(),
Port: raddr.Port,
RelAddr: RelAddr,
RelPort: RelPort,
OnClose: func() error {
client.Close()
return locConn.Close()
},
}
candidate, err := NewCandidateRelay(&relayConfig)
if err != nil {
if relayConErr := relayConn.Close(); relayConErr != nil {
a.log.Warnf("Failed to close relay %v", relayConErr)
}
client.Close()
closeConnAndLog(locConn, a.log, fmt.Sprintf("Failed to create relay candidate: %s %s: %v\n", network, raddr.String(), err))
continue
}
if err := a.addCandidate(candidate, relayConn); err != nil {
if closeErr := candidate.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)
}
}(*urls[i])
}
// Block until all STUN URLs have been gathered (or timed out)
wg.Wait()
return nil
}
+156
View File
@@ -3,12 +3,16 @@
package ice
import (
"context"
"net"
"reflect"
"sort"
"strconv"
"testing"
"time"
"github.com/pion/transport/test"
"github.com/pion/turn/v2"
"github.com/stretchr/testify/assert"
)
@@ -70,3 +74,155 @@ func TestListenUDP(t *testing.T) {
assert.NoError(t, a.Close())
}
// Assert that STUN gathering is done concurrently
func TestSTUNConcurrency(t *testing.T) {
lim := test.TimeOut(time.Second * 30)
defer lim.Stop()
report := test.CheckRoutines(t)
defer report()
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: SchemeTypeSTUN,
Host: "127.0.0.1",
Port: serverPort,
})
a, err := NewAgent(&AgentConfig{
NetworkTypes: supportedNetworkTypes,
Trickle: true,
Urls: urls,
CandidateTypes: []CandidateType{CandidateTypeServerReflexive},
})
assert.NoError(t, err)
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()
assert.NoError(t, a.Close())
assert.NoError(t, server.Close())
}
// Assert that TURN gathering is done concurrently
func TestTURNConcurrency(t *testing.T) {
lim := test.TimeOut(time.Second * 30)
defer lim.Stop()
report := test.CheckRoutines(t)
defer report()
runTest := func(protocol ProtoType, scheme SchemeType, packetConn net.PacketConn, listener net.Listener, serverPort int) {
packetConnConfigs := []turn.PacketConnConfig{}
if packetConn != nil {
packetConnConfigs = append(packetConnConfigs, turn.PacketConnConfig{
PacketConn: packetConn,
RelayAddressGenerator: &turn.RelayAddressGeneratorNone{Address: "127.0.0.1"},
})
}
listenerConfigs := []turn.ListenerConfig{}
if listener != nil {
listenerConfigs = append(listenerConfigs, turn.ListenerConfig{
Listener: listener,
RelayAddressGenerator: &turn.RelayAddressGeneratorNone{Address: "127.0.0.1"},
})
}
server, err := turn.NewServer(turn.ServerConfig{
Realm: "pion.ly",
AuthHandler: optimisticAuthHandler,
PacketConnConfigs: packetConnConfigs,
ListenerConfigs: listenerConfigs,
})
assert.NoError(t, err)
urls := []*URL{}
for i := 0; i <= 10; i++ {
urls = append(urls, &URL{
Scheme: scheme,
Host: "127.0.0.1",
Username: "username",
Password: "password",
Proto: protocol,
Port: serverPort + 1,
})
}
urls = append(urls, &URL{
Scheme: scheme,
Host: "127.0.0.1",
Username: "username",
Password: "password",
Proto: protocol,
Port: serverPort,
})
a, err := NewAgent(&AgentConfig{
NetworkTypes: supportedNetworkTypes,
Trickle: true,
Urls: urls,
CandidateTypes: []CandidateType{CandidateTypeRelay},
})
assert.NoError(t, err)
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()
assert.NoError(t, a.Close())
assert.NoError(t, server.Close())
}
t.Run("UDP Relay", func(t *testing.T) {
serverPort := randomPort(t)
serverListener, err := net.ListenPacket("udp", "127.0.0.1:"+strconv.Itoa(serverPort))
assert.NoError(t, err)
runTest(ProtoTypeUDP, SchemeTypeTURN, serverListener, nil, serverPort)
})
t.Run("TURN Relay", func(t *testing.T) {
serverPort := randomPort(t)
serverListener, err := net.Listen("tcp", "127.0.0.1:"+strconv.Itoa(serverPort))
assert.NoError(t, err)
runTest(ProtoTypeTCP, SchemeTypeTURN, nil, serverListener, serverPort)
})
}