diff --git a/README.md b/README.md index c4733ea..81c2bf5 100644 --- a/README.md +++ b/README.md @@ -45,6 +45,7 @@ Check out the **[contributing wiki](https://github.com/pion/webrtc/wiki/Contribu * [Zizheng Tai](https://github.com/ZizhengTai) * [Aaron France](https://github.com/AeroNotix) * [Chao Yuan](https://github.com/yuanchao0310) +* [Jason Maldonis](https://github.com/jjmaldonis) ### License MIT License - see [LICENSE](LICENSE) for full text diff --git a/agent.go b/agent.go index dfa0876..6b87693 100644 --- a/agent.go +++ b/agent.go @@ -51,8 +51,8 @@ const ( // the number of bytes that can be buffered before we start to error maxBufferSize = 1000 * 1000 // 1MB - // the number of outbound binding requests we cache - maxPendingBindingRequests = 50 + // wait time before binding requests can be deleted + maxBindingRequestTimeout = 500 * time.Millisecond ) var ( @@ -60,6 +60,7 @@ var ( ) type bindingRequest struct { + timestamp time.Time transactionID [stun.TransactionIDSize]byte destination net.Addr isUseCandidate bool @@ -338,7 +339,7 @@ func NewAgent(config *AgentConfig) (*Agent, error) { connectionState: ConnectionStateNew, localCandidates: make(map[NetworkType][]Candidate), remoteCandidates: make(map[NetworkType][]Candidate), - pendingBindingRequests: make([]bindingRequest, 0, maxPendingBindingRequests), + pendingBindingRequests: make([]bindingRequest, 0), checklist: make([]*candidatePair, 0), urls: config.Urls, networkTypes: config.NetworkTypes, @@ -953,17 +954,12 @@ func (a *Agent) findRemoteCandidate(networkType NetworkType, addr net.Addr) Cand func (a *Agent) sendBindingRequest(m *stun.Message, local, remote Candidate) { a.log.Tracef("ping STUN from %s to %s\n", local.String(), remote.String()) - if overflow := len(a.pendingBindingRequests) - (maxPendingBindingRequests - 1); overflow > 0 { - a.log.Debugf("Discarded %d pending binding requests, pendingBindingRequests is full", overflow) - a.pendingBindingRequests = a.pendingBindingRequests[overflow:] - } - - useCandidate := m.Contains(stun.AttrUseCandidate) - + a.invalidatePendingBindingRequests(time.Now()) a.pendingBindingRequests = append(a.pendingBindingRequests, bindingRequest{ + timestamp: time.Now(), transactionID: m.TransactionID, destination: remote.addr(), - isUseCandidate: useCandidate, + isUseCandidate: m.Contains(stun.AttrUseCandidate), }) a.sendSTUN(m, local, remote) @@ -985,9 +981,32 @@ func (a *Agent) sendBindingSuccess(m *stun.Message, local, remote Candidate) { } } +/* Removes pending binding requests that are over maxBindingRequestTimeout old + + Let HTO be the transaction timeout, which SHOULD be 2*RTT if + RTT is known or 500 ms otherwise. + https://tools.ietf.org/html/rfc8445#appendix-B.1 +*/ +func (a *Agent) invalidatePendingBindingRequests(filterTime time.Time) { + initialSize := len(a.pendingBindingRequests) + + temp := a.pendingBindingRequests[:0] + for _, bindingRequest := range a.pendingBindingRequests { + if filterTime.Sub(bindingRequest.timestamp) < maxBindingRequestTimeout { + temp = append(temp, bindingRequest) + } + } + + a.pendingBindingRequests = temp + if bindRequestsRemoved := initialSize - len(a.pendingBindingRequests); bindRequestsRemoved > 0 { + a.log.Tracef("Discarded %d binding requests because they expired", bindRequestsRemoved) + } +} + // Assert that the passed TransactionID is in our pendingBindingRequests and returns the destination // If the bindingRequest was valid remove it from our pending cache func (a *Agent) handleInboundBindingSuccess(id [stun.TransactionIDSize]byte) (bool, *bindingRequest) { + a.invalidatePendingBindingRequests(time.Now()) for i := range a.pendingBindingRequests { if a.pendingBindingRequests[i].transactionID == id { validBindingRequest := a.pendingBindingRequests[i] diff --git a/agent_test.go b/agent_test.go index 6f7e524..5a1b7ea 100644 --- a/agent_test.go +++ b/agent_test.go @@ -10,6 +10,7 @@ import ( "github.com/pion/stun" "github.com/pion/transport/test" "github.com/pion/transport/vnet" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -325,7 +326,7 @@ func TestHandlePeerReflexive(t *testing.T) { tID := [stun.TransactionIDSize]byte{} copy(tID[:], []byte("ABC")) a.pendingBindingRequests = []bindingRequest{ - {tID, &net.UDPAddr{}, false}, + {time.Now(), tID, &net.UDPAddr{}, false}, } hostConfig := CandidateHostConfig{ @@ -551,6 +552,7 @@ func TestInboundValidity(t *testing.T) { if len(a.remoteCandidates) == 1 { t.Fatal("Binding with invalid MessageIntegrity was able to create prflx candidate") } + }) t.Run("Invalid Binding success responses should be discarded", func(t *testing.T) { @@ -1200,3 +1202,27 @@ func TestInitExtIPMapping(t *testing.T) { t.Fatalf("Unexpected error: %v", err) } } + +func TestBindingRequestTimeout(t *testing.T) { + const expectedRemovalCount = 2 + + a, err := NewAgent(&AgentConfig{}) + assert.NoError(t, err) + + now := time.Now() + a.pendingBindingRequests = append(a.pendingBindingRequests, bindingRequest{ + timestamp: now, + }) + a.pendingBindingRequests = append(a.pendingBindingRequests, bindingRequest{ + timestamp: now.Add(-25 * time.Millisecond), + }) + a.pendingBindingRequests = append(a.pendingBindingRequests, bindingRequest{ + timestamp: now.Add(-750 * time.Millisecond), + }) + a.pendingBindingRequests = append(a.pendingBindingRequests, bindingRequest{ + timestamp: now.Add(-75 * time.Hour), + }) + + a.invalidatePendingBindingRequests(now) + assert.Equal(t, len(a.pendingBindingRequests), expectedRemovalCount, "Binding invalidation due to timeout did not remove the correct number of binding requests") +}