From f1bfc8fea1da22a144787916d215b8eb79561d40 Mon Sep 17 00:00:00 2001 From: Sean DuBois Date: Sat, 29 Feb 2020 16:19:03 -0800 Subject: [PATCH] Start Agent performance refactor Remove taskChan and make .run just take an Agent wide mutex and run the function. These is now a blocking operation so all channels used to communicate from it must be buffered. After this we will slowly remove usage of .run and make things more thread safe. Relates to #80, #67, #2 --- agent.go | 224 +++++++++++++++++----------------------------- candidate_base.go | 10 +-- candidatetype.go | 12 +++ gather.go | 1 - mdns.go | 39 ++++++++ selection.go | 10 +-- transport.go | 4 +- transport_test.go | 2 +- util.go | 11 +++ 9 files changed, 155 insertions(+), 158 deletions(-) diff --git a/agent.go b/agent.go index d2dff26..360af5c 100644 --- a/agent.go +++ b/agent.go @@ -16,7 +16,6 @@ import ( "github.com/pion/stun" "github.com/pion/transport/packetio" "github.com/pion/transport/vnet" - "golang.org/x/net/ipv4" ) const ( @@ -67,6 +66,10 @@ type bindingRequest struct { // Agent represents the ICE agent type Agent struct { + // Lock for transactional operations on Agent. Unlike a mutex + // all queued lock attempts are canceled when .Close() is called + muChan chan struct{} + onConnectionStateChangeHdlr func(ConnectionState) onSelectedCandidatePairChangeHdlr func(Candidate, Candidate) onCandidateHdlr func(Candidate) @@ -75,7 +78,6 @@ type Agent struct { opened bool // State owned by the taskLoop - taskChan chan task onConnected chan struct{} onConnectedOnce sync.Once @@ -132,8 +134,7 @@ type Agent struct { checklist []*candidatePair selector pairCandidateSelector - selectedPairMutex sync.RWMutex - selectedPair *candidatePair + selectedPair atomic.Value // *candidatePair urls []*URL networkTypes []NetworkType @@ -170,13 +171,30 @@ func (a *Agent) ok() error { } func (a *Agent) getErr() error { - err := a.err.Load() - if err != nil { + if err := a.err.Load(); err != nil { return err } return ErrClosed } +// Run an operation with the the lock taken +// If the agent is closed return an error +func (a *Agent) run(t func(*Agent)) error { + if err := a.ok(); err != nil { + return err + } + + select { + case <-a.done: + return a.getErr() + case a.muChan <- struct{}{}: + t(a) + + <-a.muChan + return nil + } +} + // AgentConfig collects the arguments to ice.Agent construction into // a single structure, for future-proofness of the interface type AgentConfig struct { @@ -276,49 +294,6 @@ type AgentConfig struct { InsecureSkipVerify bool } -func containsCandidateType(candidateType CandidateType, candidateTypeList []CandidateType) bool { - if candidateTypeList == nil { - return false - } - for _, ct := range candidateTypeList { - if ct == candidateType { - return true - } - } - return false -} - -func createMulticastDNS(mDNSMode MulticastDNSMode, mDNSName string, log logging.LeveledLogger) (*mdns.Conn, MulticastDNSMode, error) { - if mDNSMode == MulticastDNSModeDisabled { - return nil, mDNSMode, nil - } - - addr, mdnsErr := net.ResolveUDPAddr("udp4", mdns.DefaultAddress) - if mdnsErr != nil { - return nil, mDNSMode, mdnsErr - } - - l, mdnsErr := net.ListenUDP("udp4", addr) - if mdnsErr != nil { - // If ICE fails to start MulticastDNS server just warn the user and continue - log.Errorf("Failed to enable mDNS, continuing in mDNS disabled mode: (%s)", mdnsErr) - return nil, MulticastDNSModeDisabled, nil - } - - switch mDNSMode { - case MulticastDNSModeQueryOnly: - conn, err := mdns.Server(ipv4.NewPacketConn(l), &mdns.Config{}) - return conn, mDNSMode, err - case MulticastDNSModeQueryAndGather: - conn, err := mdns.Server(ipv4.NewPacketConn(l), &mdns.Config{ - LocalNames: []string{mDNSName}, - }) - return conn, mDNSMode, err - default: - return nil, mDNSMode, nil - } -} - // NewAgent creates a new Agent func NewAgent(config *AgentConfig) (*Agent, error) { var err error @@ -394,7 +369,6 @@ func NewAgent(config *AgentConfig) (*Agent, error) { networkTypes: config.NetworkTypes, localUfrag: localUfrag, localPwd: localPwd, - taskChan: make(chan task), onConnected: make(chan struct{}), buffer: packetio.NewBuffer(), done: make(chan struct{}), @@ -404,6 +378,7 @@ func NewAgent(config *AgentConfig) (*Agent, error) { loggerFactory: loggerFactory, log: log, net: config.Net, + muChan: make(chan struct{}, 1), mDNSMode: mDNSMode, mDNSName: mDNSName, @@ -448,8 +423,6 @@ func NewAgent(config *AgentConfig) (*Agent, error) { return nil, err } - go a.taskLoop() - // Initialize local candidates if !a.trickle { a.gatherCandidates() @@ -625,6 +598,27 @@ func (a *Agent) startConnectivityChecks(isControlling bool, remoteUfrag, remoteP // TODO this should be dynamic, and grow when the connection is stable a.requestConnectivityCheck() agent.connectivityTicker = time.NewTicker(a.taskLoopInterval) + + go func() { + contact := func() { + if err := a.run(func(a *Agent) { + a.selector.ContactCandidates() + }); err != nil { + a.log.Warnf("taskLoop failed: %v", err) + } + } + + for { + select { + case <-a.forceCandidateContact: + contact() + case <-a.connectivityTicker.C: + contact() + case <-a.done: + return + } + } + }() }) } @@ -646,10 +640,14 @@ func (a *Agent) setSelectedPair(p *candidatePair) { // Notify when the selected pair changes a.onSelectedCandidatePairChange(p) - a.selectedPairMutex.Lock() - a.selectedPair = p - a.selectedPair.nominated = true - a.selectedPairMutex.Unlock() + if p != nil { + p.nominated = true + a.selectedPair.Store(p) + } else { + var nilPair *candidatePair + a.selectedPair.Store(nilPair) + } + a.updateConnectionState(ConnectionStateConnected) // Close mDNS Conn. We don't need to do anymore querying @@ -731,65 +729,17 @@ func (a *Agent) findPair(local, remote Candidate) *candidatePair { return nil } -// A task is a -type task func(*Agent) - -func (a *Agent) run(t task) error { - err := a.ok() - if err != nil { - return err - } - - select { - case <-a.done: - return a.getErr() - case a.taskChan <- t: - } - return nil -} - -func (a *Agent) taskLoop() { - for { - if a.selector != nil { - select { - case <-a.forceCandidateContact: - a.selector.ContactCandidates() - case <-a.connectivityTicker.C: - a.selector.ContactCandidates() - case t := <-a.taskChan: - // Run the task - t(a) - - case <-a.done: - return - } - } else { - select { - case <-a.forceCandidateContact: - case t := <-a.taskChan: - // Run the task - t(a) - - case <-a.done: - return - } - } - } -} - // validateSelectedPair checks if the selected pair is (still) valid // Note: the caller should hold the agent lock. func (a *Agent) validateSelectedPair() bool { - selectedPair, err := a.getSelectedPair() - if err != nil { + selectedPair := a.getSelectedPair() + if selectedPair == nil { return false } if (a.connectionTimeout != 0) && (time.Since(selectedPair.remote.LastReceived()) > a.connectionTimeout) { - a.selectedPairMutex.Lock() - a.selectedPair = nil - a.selectedPairMutex.Unlock() + a.setSelectedPair(nil) a.updateConnectionState(ConnectionStateDisconnected) return false } @@ -801,8 +751,8 @@ func (a *Agent) validateSelectedPair() bool { // if no packet has been sent on that pair in the last keepaliveInterval // Note: the caller should hold the agent lock. func (a *Agent) checkKeepalive() { - selectedPair, err := a.getSelectedPair() - if err != nil { + selectedPair := a.getSelectedPair() + if selectedPair == nil { return } @@ -925,7 +875,7 @@ func (a *Agent) addCandidate(c Candidate, candidateConn net.PacketConn) error { // GetLocalCandidates returns the local candidates func (a *Agent) GetLocalCandidates() ([]Candidate, error) { - res := make(chan []Candidate) + res := make(chan []Candidate, 1) err := a.run(func(agent *Agent) { var candidates []Candidate @@ -984,30 +934,20 @@ func (a *Agent) Close() error { } a.closeMulticastConn() + a.updateConnectionState(ConnectionStateClosed) }) if err != nil { return err } <-done - a.updateConnectionState(ConnectionStateClosed) - return nil } func (a *Agent) findRemoteCandidate(networkType NetworkType, addr net.Addr) Candidate { - var ip net.IP - var port int - - switch casted := addr.(type) { - case *net.UDPAddr: - ip = casted.IP - port = casted.Port - case *net.TCPAddr: - ip = casted.IP - port = casted.Port - default: - a.log.Warnf("unsupported address type %T", a) + ip, port, err := addrIPAndPort(addr) + if err != nil { + a.log.Warn(err.Error()) return nil } @@ -1175,28 +1115,30 @@ func (a *Agent) handleInbound(m *stun.Message, local Candidate, remote net.Addr) } } -// noSTUNSeen processes non STUN traffic from a remote candidate, +// validateNonSTUNTraffic processes non STUN traffic from a remote candidate, // and returns true if it is an actual remote candidate -func (a *Agent) noSTUNSeen(local Candidate, remote net.Addr) bool { - remoteCandidate := a.findRemoteCandidate(local.NetworkType(), remote) - if remoteCandidate == nil { - return false +func (a *Agent) validateNonSTUNTraffic(local Candidate, remote net.Addr) bool { + var isValidCandidate uint64 + if err := a.run(func(agent *Agent) { + remoteCandidate := a.findRemoteCandidate(local.NetworkType(), remote) + if remoteCandidate != nil { + remoteCandidate.seen(false) + atomic.AddUint64(&isValidCandidate, 1) + } + }); err != nil { + a.log.Warnf("failed to validate remote candidate: %v", err) } - remoteCandidate.seen(false) - return true + return atomic.LoadUint64(&isValidCandidate) == 1 } -func (a *Agent) getSelectedPair() (*candidatePair, error) { - a.selectedPairMutex.RLock() - selectedPair := a.selectedPair - a.selectedPairMutex.RUnlock() - +func (a *Agent) getSelectedPair() *candidatePair { + selectedPair := a.selectedPair.Load() if selectedPair == nil { - return nil, ErrNoCandidatePairs + return nil } - return selectedPair, nil + return selectedPair.(*candidatePair) } func (a *Agent) closeMulticastConn() { @@ -1209,7 +1151,7 @@ func (a *Agent) closeMulticastConn() { // GetCandidatePairsStats returns a list of candidate pair stats func (a *Agent) GetCandidatePairsStats() []CandidatePairStats { - resultChan := make(chan []CandidatePairStats) + resultChan := make(chan []CandidatePairStats, 1) err := a.run(func(agent *Agent) { result := make([]CandidatePairStats, 0, len(agent.checklist)) for _, cp := range agent.checklist { @@ -1255,7 +1197,7 @@ func (a *Agent) GetCandidatePairsStats() []CandidatePairStats { // GetLocalCandidatesStats returns a list of local candidates stats func (a *Agent) GetLocalCandidatesStats() []CandidateStats { - resultChan := make(chan []CandidateStats) + resultChan := make(chan []CandidateStats, 1) err := a.run(func(agent *Agent) { result := make([]CandidateStats, 0, len(agent.localCandidates)) for networkType, localCandidates := range agent.localCandidates { @@ -1286,7 +1228,7 @@ func (a *Agent) GetLocalCandidatesStats() []CandidateStats { // GetRemoteCandidatesStats returns a list of remote candidates stats func (a *Agent) GetRemoteCandidatesStats() []CandidateStats { - resultChan := make(chan []CandidateStats) + resultChan := make(chan []CandidateStats, 1) err := a.run(func(agent *Agent) { result := make([]CandidateStats, 0, len(agent.remoteCandidates)) for networkType, localCandidates := range agent.remoteCandidates { diff --git a/candidate_base.go b/candidate_base.go index bd244f9..7417ea0 100644 --- a/candidate_base.go +++ b/candidate_base.go @@ -119,15 +119,9 @@ func handleInboundCandidateMsg(c Candidate, buffer []byte, srcAddr net.Addr, log return } - isValidRemoteCandidate := make(chan bool, 1) - err := c.agent().run(func(agent *Agent) { - isValidRemoteCandidate <- agent.noSTUNSeen(c, srcAddr) - }) - - if err != nil { - log.Warnf("Failed to handle message: %v", err) - } else if !<-isValidRemoteCandidate { + if !c.agent().validateNonSTUNTraffic(c, srcAddr) { log.Warnf("Discarded message from %s, not a valid remote candidate", c.addr()) + return } // NOTE This will return packetio.ErrFull if the buffer ever manages to fill up. diff --git a/candidatetype.go b/candidatetype.go index 6adadfa..db12af8 100644 --- a/candidatetype.go +++ b/candidatetype.go @@ -44,3 +44,15 @@ func (c CandidateType) Preference() uint16 { } return 0 } + +func containsCandidateType(candidateType CandidateType, candidateTypeList []CandidateType) bool { + if candidateTypeList == nil { + return false + } + for _, ct := range candidateTypeList { + if ct == candidateType { + return true + } + } + return false +} diff --git a/gather.go b/gather.go index ece63eb..3b2638d 100644 --- a/gather.go +++ b/gather.go @@ -95,7 +95,6 @@ func (a *Agent) gatherCandidates() { } } } - if err := a.run(func(agent *Agent) { if a.onCandidateHdlr != nil { go a.onCandidateHdlr(nil) diff --git a/mdns.go b/mdns.go index 7188b13..0895355 100644 --- a/mdns.go +++ b/mdns.go @@ -1,5 +1,13 @@ package ice +import ( + "net" + + "github.com/pion/logging" + "github.com/pion/mdns" + "golang.org/x/net/ipv4" +) + // MulticastDNSMode represents the different Multicast modes ICE can run in type MulticastDNSMode byte @@ -18,3 +26,34 @@ const ( func generateMulticastDNSName() (string, error) { return generateRandString("", ".local") } + +func createMulticastDNS(mDNSMode MulticastDNSMode, mDNSName string, log logging.LeveledLogger) (*mdns.Conn, MulticastDNSMode, error) { + if mDNSMode == MulticastDNSModeDisabled { + return nil, mDNSMode, nil + } + + addr, mdnsErr := net.ResolveUDPAddr("udp4", mdns.DefaultAddress) + if mdnsErr != nil { + return nil, mDNSMode, mdnsErr + } + + l, mdnsErr := net.ListenUDP("udp4", addr) + if mdnsErr != nil { + // If ICE fails to start MulticastDNS server just warn the user and continue + log.Errorf("Failed to enable mDNS, continuing in mDNS disabled mode: (%s)", mdnsErr) + return nil, MulticastDNSModeDisabled, nil + } + + switch mDNSMode { + case MulticastDNSModeQueryOnly: + conn, err := mdns.Server(ipv4.NewPacketConn(l), &mdns.Config{}) + return conn, mDNSMode, err + case MulticastDNSModeQueryAndGather: + conn, err := mdns.Server(ipv4.NewPacketConn(l), &mdns.Config{ + LocalNames: []string{mDNSName}, + }) + return conn, mDNSMode, err + default: + return nil, mDNSMode, nil + } +} diff --git a/selection.go b/selection.go index 9cbe6c0..45a6e4e 100644 --- a/selection.go +++ b/selection.go @@ -71,7 +71,7 @@ func (s *controllingSelector) isNominatable(c Candidate) bool { func (s *controllingSelector) ContactCandidates() { switch { - case s.agent.selectedPair != nil: + case s.agent.getSelectedPair() != nil: if s.agent.validateSelectedPair() { s.log.Trace("checking keepalive") s.agent.checkKeepalive() @@ -130,7 +130,7 @@ func (s *controllingSelector) HandleBindingRequest(m *stun.Message, local, remot return } - if p.state == CandidatePairStateSucceeded && s.nominatedPair == nil && s.agent.selectedPair == nil { + if p.state == CandidatePairStateSucceeded && s.nominatedPair == nil && s.agent.getSelectedPair() == nil { bestPair := s.agent.getBestAvailableCandidatePair() if bestPair == nil { s.log.Tracef("No best pair available\n") @@ -170,7 +170,7 @@ func (s *controllingSelector) HandleSuccessResponse(m *stun.Message, local, remo p.state = CandidatePairStateSucceeded s.log.Tracef("Found valid candidate pair: %s", p) - if pendingRequest.isUseCandidate && s.agent.selectedPair == nil { + if pendingRequest.isUseCandidate && s.agent.getSelectedPair() == nil { s.agent.setSelectedPair(p) } } @@ -203,7 +203,7 @@ func (s *controlledSelector) Start() { } func (s *controlledSelector) ContactCandidates() { - if s.agent.selectedPair != nil { + if s.agent.getSelectedPair() != nil { if s.agent.validateSelectedPair() { s.log.Trace("checking keepalive") s.agent.checkKeepalive() @@ -288,7 +288,7 @@ 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 s.agent.selectedPair == nil { + if selectedPair := s.agent.getSelectedPair(); selectedPair == nil { s.agent.setSelectedPair(p) } s.agent.sendBindingSuccess(m, local, remote) diff --git a/transport.go b/transport.go index 9468262..9641f45 100644 --- a/transport.go +++ b/transport.go @@ -91,8 +91,8 @@ func (c *Conn) Write(p []byte) (int, error) { return 0, errors.New("the ICE conn can't write STUN messages") } - pair, err := c.agent.getSelectedPair() - if err != nil { + pair := c.agent.getSelectedPair() + if pair == nil { return 0, err } diff --git a/transport_test.go b/transport_test.go index 5358ca0..97e3f4f 100644 --- a/transport_test.go +++ b/transport_test.go @@ -29,7 +29,7 @@ func TestStressDuplex(t *testing.T) { func testTimeout(t *testing.T, c *Conn, timeout time.Duration) { const pollrate = 100 * time.Millisecond const margin = 20 * time.Millisecond // allow 20msec error in time - statechan := make(chan ConnectionState) + statechan := make(chan ConnectionState, 1) ticker := time.NewTicker(pollrate) startedAt := time.Now() diff --git a/util.go b/util.go index fae3b87..d39be68 100644 --- a/util.go +++ b/util.go @@ -246,3 +246,14 @@ func listenUDPInPortRange(vnet *vnet.Net, log logging.LeveledLogger, portMax, po } return nil, ErrPort } + +func addrIPAndPort(addr net.Addr) (net.IP, int, error) { + switch casted := addr.(type) { + case *net.UDPAddr: + return casted.IP, casted.Port, nil + case *net.TCPAddr: + return casted.IP, casted.Port, nil + default: + return nil, 0, fmt.Errorf("unsupported address type %T", addr) + } +}