diff --git a/pkg/rtc/mediatracksubscriptions.go b/pkg/rtc/mediatracksubscriptions.go index a34271180..92a88e78f 100644 --- a/pkg/rtc/mediatracksubscriptions.go +++ b/pkg/rtc/mediatracksubscriptions.go @@ -10,6 +10,7 @@ import ( "github.com/pion/rtcp" "github.com/pion/webrtc/v3" "github.com/pion/webrtc/v3/pkg/rtcerr" + "go.uber.org/atomic" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" @@ -171,8 +172,13 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * }) // Bind callback can happen from replaceTrack, so set it up early + var reusingTransceiver atomic.Bool + var forwarderState sfu.ForwarderState downTrack.OnBind(func() { wr.DetermineReceiver(downTrack.Codec()) + if reusingTransceiver.Load() { + downTrack.SeedForwarderState(forwarderState) + } if err = wr.AddDownTrack(downTrack); err != nil { t.params.Logger.Errorw("could not add down track", err, "participant", sub.Identity(), "pID", sub.ID()) } @@ -191,9 +197,11 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * var sender *webrtc.RTPSender // try cached RTP senders for a chance to replace track + var existingTransceiver *webrtc.RTPTransceiver replacedTrack := false - existingTransceiver := sub.GetCachedRTPTransceiver(trackID) + existingTransceiver, forwarderState = sub.GetCachedDownTrack(trackID) if existingTransceiver != nil { + reusingTransceiver.Store(true) rtpSender := existingTransceiver.Sender() if rtpSender != nil { err := rtpSender.ReplaceTrack(downTrack) @@ -216,6 +224,7 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * existingTransceiver.Stop() } } + reusingTransceiver.Store(false) // if cannot replace, find an unused transceiver or add new one if transceiver == nil { @@ -259,7 +268,7 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * // wthether re-using or stopping remove transceiver from cache // NOTE: safety net, if somehow a cached transceiver is re-used by a different track - sub.UncacheRTPTransceiver(transceiver) + sub.UncacheDownTrack(transceiver) sendParameters := sender.GetParameters() downTrack.SetRTPHeaderExtensions(sendParameters.HeaderExtensions) @@ -339,15 +348,15 @@ func (t *MediaTrackSubscriptions) closeSubscribedTrack(subTrack types.Subscribed return } + dt.CloseWithFlush(!willBeResumed) + if willBeResumed { tr := dt.GetTransceiver() if tr != nil { sub := subTrack.Subscriber() - sub.CacheRTPTransceiver(subTrack.ID(), tr) + sub.CacheDownTrack(subTrack.ID(), tr, dt.GetForwarderState()) } } - - dt.CloseWithFlush(!willBeResumed) } func (t *MediaTrackSubscriptions) ResyncAllSubscribers() { diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 9c02eaea7..ba2696ea8 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -44,6 +44,11 @@ type pendingTrackInfo struct { migrated bool } +type downTrackState struct { + transceiver *webrtc.RTPTransceiver + forwarder sfu.ForwarderState +} + type ParticipantParams struct { Identity livekit.ParticipantIdentity Name livekit.ParticipantName @@ -138,7 +143,7 @@ type ParticipantImpl struct { firstConnected atomic.Bool iceConfig types.IceConfig - cachedRTPTransceivers map[livekit.TrackID]*webrtc.RTPTransceiver + cachedDownTracks map[livekit.TrackID]*downTrackState } func NewParticipant(params ParticipantParams) (*ParticipantImpl, error) { @@ -161,7 +166,7 @@ func NewParticipant(params ParticipantParams) (*ParticipantImpl, error) { subscribedTo: make(map[livekit.ParticipantID]struct{}), connectedAt: time.Now(), rttUpdatedAt: time.Now(), - cachedRTPTransceivers: make(map[livekit.TrackID]*webrtc.RTPTransceiver), + cachedDownTracks: make(map[livekit.TrackID]*downTrackState), } p.version.Store(params.InitialVersion) p.migrateState.Store(types.MigrateStateInit) @@ -1912,29 +1917,34 @@ func (p *ParticipantImpl) setDowntracksConnected() { } } -func (p *ParticipantImpl) CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver) { +func (p *ParticipantImpl) CacheDownTrack(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver, forwarderState sfu.ForwarderState) { p.lock.Lock() - if existing := p.cachedRTPTransceivers[trackID]; existing != nil && existing != rtpTransceiver { + if existing := p.cachedDownTracks[trackID]; existing != nil && existing.transceiver != rtpTransceiver { p.params.Logger.Infow("cached transceiver change", "trackID", trackID) } - p.cachedRTPTransceivers[trackID] = rtpTransceiver + p.cachedDownTracks[trackID] = &downTrackState{transceiver: rtpTransceiver, forwarder: forwarderState} p.lock.Unlock() } -func (p *ParticipantImpl) UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver) { +func (p *ParticipantImpl) UncacheDownTrack(rtpTransceiver *webrtc.RTPTransceiver) { p.lock.Lock() - for trackID, tr := range p.cachedRTPTransceivers { - if tr == rtpTransceiver { - delete(p.cachedRTPTransceivers, trackID) + for trackID, dts := range p.cachedDownTracks { + if dts.transceiver == rtpTransceiver { + delete(p.cachedDownTracks, trackID) break } } p.lock.Unlock() } -func (p *ParticipantImpl) GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver { +func (p *ParticipantImpl) GetCachedDownTrack(trackID livekit.TrackID) (*webrtc.RTPTransceiver, sfu.ForwarderState) { p.lock.RLock() defer p.lock.RUnlock() - return p.cachedRTPTransceivers[trackID] + dts := p.cachedDownTracks[trackID] + if dts != nil { + return dts.transceiver, dts.forwarder + } + + return nil, sfu.ForwarderState{} } diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index a0710a81b..0133e748a 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -247,9 +247,9 @@ type LocalParticipant interface { UpdateRTT(rtt uint32) - CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver) - UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver) - GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver + CacheDownTrack(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver, forwarderState sfu.ForwarderState) + UncacheDownTrack(rtpTransceiver *webrtc.RTPTransceiver) + GetCachedDownTrack(trackID livekit.TrackID) (*webrtc.RTPTransceiver, sfu.ForwarderState) } // Room is a container of participants, and can provide room-level actions diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index 219936036..cc0590016 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -7,6 +7,7 @@ import ( "github.com/livekit/livekit-server/pkg/routing" "github.com/livekit/livekit-server/pkg/rtc/types" + "github.com/livekit/livekit-server/pkg/sfu" "github.com/livekit/protocol/auth" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" @@ -50,11 +51,12 @@ type FakeLocalParticipant struct { addTrackArgsForCall []struct { arg1 *livekit.AddTrackRequest } - CacheRTPTransceiverStub func(livekit.TrackID, *webrtc.RTPTransceiver) - cacheRTPTransceiverMutex sync.RWMutex - cacheRTPTransceiverArgsForCall []struct { + CacheDownTrackStub func(livekit.TrackID, *webrtc.RTPTransceiver, sfu.ForwarderState) + cacheDownTrackMutex sync.RWMutex + cacheDownTrackArgsForCall []struct { arg1 livekit.TrackID arg2 *webrtc.RTPTransceiver + arg3 sfu.ForwarderState } CanPublishStub func() bool canPublishMutex sync.RWMutex @@ -150,16 +152,18 @@ type FakeLocalParticipant struct { result1 float64 result2 bool } - GetCachedRTPTransceiverStub func(livekit.TrackID) *webrtc.RTPTransceiver - getCachedRTPTransceiverMutex sync.RWMutex - getCachedRTPTransceiverArgsForCall []struct { + GetCachedDownTrackStub func(livekit.TrackID) (*webrtc.RTPTransceiver, sfu.ForwarderState) + getCachedDownTrackMutex sync.RWMutex + getCachedDownTrackArgsForCall []struct { arg1 livekit.TrackID } - getCachedRTPTransceiverReturns struct { + getCachedDownTrackReturns struct { result1 *webrtc.RTPTransceiver + result2 sfu.ForwarderState } - getCachedRTPTransceiverReturnsOnCall map[int]struct { + getCachedDownTrackReturnsOnCall map[int]struct { result1 *webrtc.RTPTransceiver + result2 sfu.ForwarderState } GetConnectionQualityStub func() *livekit.ConnectionQualityInfo getConnectionQualityMutex sync.RWMutex @@ -606,9 +610,9 @@ type FakeLocalParticipant struct { toProtoReturnsOnCall map[int]struct { result1 *livekit.ParticipantInfo } - UncacheRTPTransceiverStub func(*webrtc.RTPTransceiver) - uncacheRTPTransceiverMutex sync.RWMutex - uncacheRTPTransceiverArgsForCall []struct { + UncacheDownTrackStub func(*webrtc.RTPTransceiver) + uncacheDownTrackMutex sync.RWMutex + uncacheDownTrackArgsForCall []struct { arg1 *webrtc.RTPTransceiver } UpdateMediaLossStub func(livekit.NodeID, livekit.TrackID, uint32) error @@ -873,37 +877,38 @@ func (fake *FakeLocalParticipant) AddTrackArgsForCall(i int) *livekit.AddTrackRe return argsForCall.arg1 } -func (fake *FakeLocalParticipant) CacheRTPTransceiver(arg1 livekit.TrackID, arg2 *webrtc.RTPTransceiver) { - fake.cacheRTPTransceiverMutex.Lock() - fake.cacheRTPTransceiverArgsForCall = append(fake.cacheRTPTransceiverArgsForCall, struct { +func (fake *FakeLocalParticipant) CacheDownTrack(arg1 livekit.TrackID, arg2 *webrtc.RTPTransceiver, arg3 sfu.ForwarderState) { + fake.cacheDownTrackMutex.Lock() + fake.cacheDownTrackArgsForCall = append(fake.cacheDownTrackArgsForCall, struct { arg1 livekit.TrackID arg2 *webrtc.RTPTransceiver - }{arg1, arg2}) - stub := fake.CacheRTPTransceiverStub - fake.recordInvocation("CacheRTPTransceiver", []interface{}{arg1, arg2}) - fake.cacheRTPTransceiverMutex.Unlock() + arg3 sfu.ForwarderState + }{arg1, arg2, arg3}) + stub := fake.CacheDownTrackStub + fake.recordInvocation("CacheDownTrack", []interface{}{arg1, arg2, arg3}) + fake.cacheDownTrackMutex.Unlock() if stub != nil { - fake.CacheRTPTransceiverStub(arg1, arg2) + fake.CacheDownTrackStub(arg1, arg2, arg3) } } -func (fake *FakeLocalParticipant) CacheRTPTransceiverCallCount() int { - fake.cacheRTPTransceiverMutex.RLock() - defer fake.cacheRTPTransceiverMutex.RUnlock() - return len(fake.cacheRTPTransceiverArgsForCall) +func (fake *FakeLocalParticipant) CacheDownTrackCallCount() int { + fake.cacheDownTrackMutex.RLock() + defer fake.cacheDownTrackMutex.RUnlock() + return len(fake.cacheDownTrackArgsForCall) } -func (fake *FakeLocalParticipant) CacheRTPTransceiverCalls(stub func(livekit.TrackID, *webrtc.RTPTransceiver)) { - fake.cacheRTPTransceiverMutex.Lock() - defer fake.cacheRTPTransceiverMutex.Unlock() - fake.CacheRTPTransceiverStub = stub +func (fake *FakeLocalParticipant) CacheDownTrackCalls(stub func(livekit.TrackID, *webrtc.RTPTransceiver, sfu.ForwarderState)) { + fake.cacheDownTrackMutex.Lock() + defer fake.cacheDownTrackMutex.Unlock() + fake.CacheDownTrackStub = stub } -func (fake *FakeLocalParticipant) CacheRTPTransceiverArgsForCall(i int) (livekit.TrackID, *webrtc.RTPTransceiver) { - fake.cacheRTPTransceiverMutex.RLock() - defer fake.cacheRTPTransceiverMutex.RUnlock() - argsForCall := fake.cacheRTPTransceiverArgsForCall[i] - return argsForCall.arg1, argsForCall.arg2 +func (fake *FakeLocalParticipant) CacheDownTrackArgsForCall(i int) (livekit.TrackID, *webrtc.RTPTransceiver, sfu.ForwarderState) { + fake.cacheDownTrackMutex.RLock() + defer fake.cacheDownTrackMutex.RUnlock() + argsForCall := fake.cacheDownTrackArgsForCall[i] + return argsForCall.arg1, argsForCall.arg2, argsForCall.arg3 } func (fake *FakeLocalParticipant) CanPublish() bool { @@ -1395,65 +1400,68 @@ func (fake *FakeLocalParticipant) GetAudioLevelReturnsOnCall(i int, result1 floa }{result1, result2} } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiver(arg1 livekit.TrackID) *webrtc.RTPTransceiver { - fake.getCachedRTPTransceiverMutex.Lock() - ret, specificReturn := fake.getCachedRTPTransceiverReturnsOnCall[len(fake.getCachedRTPTransceiverArgsForCall)] - fake.getCachedRTPTransceiverArgsForCall = append(fake.getCachedRTPTransceiverArgsForCall, struct { +func (fake *FakeLocalParticipant) GetCachedDownTrack(arg1 livekit.TrackID) (*webrtc.RTPTransceiver, sfu.ForwarderState) { + fake.getCachedDownTrackMutex.Lock() + ret, specificReturn := fake.getCachedDownTrackReturnsOnCall[len(fake.getCachedDownTrackArgsForCall)] + fake.getCachedDownTrackArgsForCall = append(fake.getCachedDownTrackArgsForCall, struct { arg1 livekit.TrackID }{arg1}) - stub := fake.GetCachedRTPTransceiverStub - fakeReturns := fake.getCachedRTPTransceiverReturns - fake.recordInvocation("GetCachedRTPTransceiver", []interface{}{arg1}) - fake.getCachedRTPTransceiverMutex.Unlock() + stub := fake.GetCachedDownTrackStub + fakeReturns := fake.getCachedDownTrackReturns + fake.recordInvocation("GetCachedDownTrack", []interface{}{arg1}) + fake.getCachedDownTrackMutex.Unlock() if stub != nil { return stub(arg1) } if specificReturn { - return ret.result1 + return ret.result1, ret.result2 } - return fakeReturns.result1 + return fakeReturns.result1, fakeReturns.result2 } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCallCount() int { - fake.getCachedRTPTransceiverMutex.RLock() - defer fake.getCachedRTPTransceiverMutex.RUnlock() - return len(fake.getCachedRTPTransceiverArgsForCall) +func (fake *FakeLocalParticipant) GetCachedDownTrackCallCount() int { + fake.getCachedDownTrackMutex.RLock() + defer fake.getCachedDownTrackMutex.RUnlock() + return len(fake.getCachedDownTrackArgsForCall) } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCalls(stub func(livekit.TrackID) *webrtc.RTPTransceiver) { - fake.getCachedRTPTransceiverMutex.Lock() - defer fake.getCachedRTPTransceiverMutex.Unlock() - fake.GetCachedRTPTransceiverStub = stub +func (fake *FakeLocalParticipant) GetCachedDownTrackCalls(stub func(livekit.TrackID) (*webrtc.RTPTransceiver, sfu.ForwarderState)) { + fake.getCachedDownTrackMutex.Lock() + defer fake.getCachedDownTrackMutex.Unlock() + fake.GetCachedDownTrackStub = stub } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiverArgsForCall(i int) livekit.TrackID { - fake.getCachedRTPTransceiverMutex.RLock() - defer fake.getCachedRTPTransceiverMutex.RUnlock() - argsForCall := fake.getCachedRTPTransceiverArgsForCall[i] +func (fake *FakeLocalParticipant) GetCachedDownTrackArgsForCall(i int) livekit.TrackID { + fake.getCachedDownTrackMutex.RLock() + defer fake.getCachedDownTrackMutex.RUnlock() + argsForCall := fake.getCachedDownTrackArgsForCall[i] return argsForCall.arg1 } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiverReturns(result1 *webrtc.RTPTransceiver) { - fake.getCachedRTPTransceiverMutex.Lock() - defer fake.getCachedRTPTransceiverMutex.Unlock() - fake.GetCachedRTPTransceiverStub = nil - fake.getCachedRTPTransceiverReturns = struct { +func (fake *FakeLocalParticipant) GetCachedDownTrackReturns(result1 *webrtc.RTPTransceiver, result2 sfu.ForwarderState) { + fake.getCachedDownTrackMutex.Lock() + defer fake.getCachedDownTrackMutex.Unlock() + fake.GetCachedDownTrackStub = nil + fake.getCachedDownTrackReturns = struct { result1 *webrtc.RTPTransceiver - }{result1} + result2 sfu.ForwarderState + }{result1, result2} } -func (fake *FakeLocalParticipant) GetCachedRTPTransceiverReturnsOnCall(i int, result1 *webrtc.RTPTransceiver) { - fake.getCachedRTPTransceiverMutex.Lock() - defer fake.getCachedRTPTransceiverMutex.Unlock() - fake.GetCachedRTPTransceiverStub = nil - if fake.getCachedRTPTransceiverReturnsOnCall == nil { - fake.getCachedRTPTransceiverReturnsOnCall = make(map[int]struct { +func (fake *FakeLocalParticipant) GetCachedDownTrackReturnsOnCall(i int, result1 *webrtc.RTPTransceiver, result2 sfu.ForwarderState) { + fake.getCachedDownTrackMutex.Lock() + defer fake.getCachedDownTrackMutex.Unlock() + fake.GetCachedDownTrackStub = nil + if fake.getCachedDownTrackReturnsOnCall == nil { + fake.getCachedDownTrackReturnsOnCall = make(map[int]struct { result1 *webrtc.RTPTransceiver + result2 sfu.ForwarderState }) } - fake.getCachedRTPTransceiverReturnsOnCall[i] = struct { + fake.getCachedDownTrackReturnsOnCall[i] = struct { result1 *webrtc.RTPTransceiver - }{result1} + result2 sfu.ForwarderState + }{result1, result2} } func (fake *FakeLocalParticipant) GetConnectionQuality() *livekit.ConnectionQualityInfo { @@ -3921,35 +3929,35 @@ func (fake *FakeLocalParticipant) ToProtoReturnsOnCall(i int, result1 *livekit.P }{result1} } -func (fake *FakeLocalParticipant) UncacheRTPTransceiver(arg1 *webrtc.RTPTransceiver) { - fake.uncacheRTPTransceiverMutex.Lock() - fake.uncacheRTPTransceiverArgsForCall = append(fake.uncacheRTPTransceiverArgsForCall, struct { +func (fake *FakeLocalParticipant) UncacheDownTrack(arg1 *webrtc.RTPTransceiver) { + fake.uncacheDownTrackMutex.Lock() + fake.uncacheDownTrackArgsForCall = append(fake.uncacheDownTrackArgsForCall, struct { arg1 *webrtc.RTPTransceiver }{arg1}) - stub := fake.UncacheRTPTransceiverStub - fake.recordInvocation("UncacheRTPTransceiver", []interface{}{arg1}) - fake.uncacheRTPTransceiverMutex.Unlock() + stub := fake.UncacheDownTrackStub + fake.recordInvocation("UncacheDownTrack", []interface{}{arg1}) + fake.uncacheDownTrackMutex.Unlock() if stub != nil { - fake.UncacheRTPTransceiverStub(arg1) + fake.UncacheDownTrackStub(arg1) } } -func (fake *FakeLocalParticipant) UncacheRTPTransceiverCallCount() int { - fake.uncacheRTPTransceiverMutex.RLock() - defer fake.uncacheRTPTransceiverMutex.RUnlock() - return len(fake.uncacheRTPTransceiverArgsForCall) +func (fake *FakeLocalParticipant) UncacheDownTrackCallCount() int { + fake.uncacheDownTrackMutex.RLock() + defer fake.uncacheDownTrackMutex.RUnlock() + return len(fake.uncacheDownTrackArgsForCall) } -func (fake *FakeLocalParticipant) UncacheRTPTransceiverCalls(stub func(*webrtc.RTPTransceiver)) { - fake.uncacheRTPTransceiverMutex.Lock() - defer fake.uncacheRTPTransceiverMutex.Unlock() - fake.UncacheRTPTransceiverStub = stub +func (fake *FakeLocalParticipant) UncacheDownTrackCalls(stub func(*webrtc.RTPTransceiver)) { + fake.uncacheDownTrackMutex.Lock() + defer fake.uncacheDownTrackMutex.Unlock() + fake.UncacheDownTrackStub = stub } -func (fake *FakeLocalParticipant) UncacheRTPTransceiverArgsForCall(i int) *webrtc.RTPTransceiver { - fake.uncacheRTPTransceiverMutex.RLock() - defer fake.uncacheRTPTransceiverMutex.RUnlock() - argsForCall := fake.uncacheRTPTransceiverArgsForCall[i] +func (fake *FakeLocalParticipant) UncacheDownTrackArgsForCall(i int) *webrtc.RTPTransceiver { + fake.uncacheDownTrackMutex.RLock() + defer fake.uncacheDownTrackMutex.RUnlock() + argsForCall := fake.uncacheDownTrackArgsForCall[i] return argsForCall.arg1 } @@ -4313,8 +4321,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.addSubscriberMutex.RUnlock() fake.addTrackMutex.RLock() defer fake.addTrackMutex.RUnlock() - fake.cacheRTPTransceiverMutex.RLock() - defer fake.cacheRTPTransceiverMutex.RUnlock() + fake.cacheDownTrackMutex.RLock() + defer fake.cacheDownTrackMutex.RUnlock() fake.canPublishMutex.RLock() defer fake.canPublishMutex.RUnlock() fake.canPublishDataMutex.RLock() @@ -4333,8 +4341,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.getAdaptiveStreamMutex.RUnlock() fake.getAudioLevelMutex.RLock() defer fake.getAudioLevelMutex.RUnlock() - fake.getCachedRTPTransceiverMutex.RLock() - defer fake.getCachedRTPTransceiverMutex.RUnlock() + fake.getCachedDownTrackMutex.RLock() + defer fake.getCachedDownTrackMutex.RUnlock() fake.getConnectionQualityMutex.RLock() defer fake.getConnectionQualityMutex.RUnlock() fake.getLoggerMutex.RLock() @@ -4437,8 +4445,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.subscriptionPermissionUpdateMutex.RUnlock() fake.toProtoMutex.RLock() defer fake.toProtoMutex.RUnlock() - fake.uncacheRTPTransceiverMutex.RLock() - defer fake.uncacheRTPTransceiverMutex.RUnlock() + fake.uncacheDownTrackMutex.RLock() + defer fake.uncacheDownTrackMutex.RUnlock() fake.updateMediaLossMutex.RLock() defer fake.updateMediaLossMutex.RUnlock() fake.updateRTTMutex.RLock() diff --git a/pkg/sfu/downtrack.go b/pkg/sfu/downtrack.go index eba45dacc..efaa92e5f 100644 --- a/pkg/sfu/downtrack.go +++ b/pkg/sfu/downtrack.go @@ -722,6 +722,14 @@ func (d *DownTrack) MaxLayers() VideoLayers { return d.forwarder.MaxLayers() } +func (d *DownTrack) GetForwarderState() ForwarderState { + return d.forwarder.GetState() +} + +func (d *DownTrack) SeedForwarderState(state ForwarderState) { + d.forwarder.SeedState(state) +} + func (d *DownTrack) GetForwardingStatus() ForwardingStatus { return d.forwarder.GetForwardingStatus() } diff --git a/pkg/sfu/forwarder.go b/pkg/sfu/forwarder.go index e270ffa4a..fb9330c68 100644 --- a/pkg/sfu/forwarder.go +++ b/pkg/sfu/forwarder.go @@ -163,6 +163,14 @@ var ( // ------------------------------------------------------------------- +type ForwarderState struct { + LastTSCalc int64 + RTP RTPMungerState + VP8 VP8MungerState +} + +// ------------------------------------------------------------------- + type Forwarder struct { lock sync.RWMutex codec webrtc.RTPCodecCapability @@ -217,6 +225,9 @@ func NewForwarder(kind webrtc.RTPCodecType, logger logger.Logger) *Forwarder { } func (f *Forwarder) DetermineCodec(codec webrtc.RTPCodecCapability) { + f.lock.Lock() + defer f.lock.Unlock() + if f.codec.MimeType != "" { return } @@ -233,6 +244,35 @@ func (f *Forwarder) DetermineCodec(codec webrtc.RTPCodecCapability) { } } +func (f *Forwarder) GetState() ForwarderState { + f.lock.RLock() + defer f.lock.RUnlock() + + state := ForwarderState{ + LastTSCalc: f.lTSCalc, + RTP: f.rtpMunger.GetLast(), + } + + if f.vp8Munger != nil { + state.VP8 = f.vp8Munger.GetLast() + } + + return state +} + +func (f *Forwarder) SeedState(state ForwarderState) { + f.lock.Lock() + defer f.lock.Unlock() + + f.lTSCalc = state.LastTSCalc + f.rtpMunger.SeedLast(state.RTP) + if f.vp8Munger != nil { + f.vp8Munger.SeedLast(state.VP8) + } + + f.started = true +} + func (f *Forwarder) Mute(val bool) (bool, VideoLayers) { f.lock.Lock() defer f.lock.Unlock() diff --git a/pkg/sfu/rtpmunger.go b/pkg/sfu/rtpmunger.go index b99c0bb9e..46b37b735 100644 --- a/pkg/sfu/rtpmunger.go +++ b/pkg/sfu/rtpmunger.go @@ -36,6 +36,11 @@ type SnTs struct { timestamp uint32 } +type RTPMungerState struct { + LastSN uint16 + LastTS uint32 +} + type RTPMungerParams struct { highestIncomingSN uint16 lastSN uint16 @@ -75,6 +80,18 @@ func (r *RTPMunger) GetParams() RTPMungerParams { } } +func (r *RTPMunger) GetLast() RTPMungerState { + return RTPMungerState{ + LastSN: r.lastSN, + LastTS: r.lastTS, + } +} + +func (r *RTPMunger) SeedLast(state RTPMungerState) { + r.lastSN = state.LastSN + r.lastTS = state.LastTS +} + func (r *RTPMunger) SetLastSnTs(extPkt *buffer.ExtPacket) { r.highestIncomingSN = extPkt.Packet.SequenceNumber - 1 r.lastSN = extPkt.Packet.SequenceNumber diff --git a/pkg/sfu/vp8munger.go b/pkg/sfu/vp8munger.go index 6197ce98d..a610a2359 100644 --- a/pkg/sfu/vp8munger.go +++ b/pkg/sfu/vp8munger.go @@ -15,6 +15,16 @@ type TranslationParamsVP8 struct { Header *buffer.VP8 } +type VP8MungerState struct { + ExtLastPictureId int32 + PictureIdUsed int + LastTl0PicIdx uint8 + Tl0PicIdxUsed int + TidUsed int + LastKeyIdx uint8 + KeyIdxUsed int +} + type VP8MungerParams struct { pictureIdWrapHandler VP8PictureIdWrapHandler extLastPictureId int32 @@ -48,6 +58,28 @@ func NewVP8Munger(logger logger.Logger) *VP8Munger { } } +func (v *VP8Munger) GetLast() VP8MungerState { + return VP8MungerState{ + ExtLastPictureId: v.extLastPictureId, + PictureIdUsed: v.pictureIdUsed, + LastTl0PicIdx: v.lastTl0PicIdx, + Tl0PicIdxUsed: v.tl0PicIdxUsed, + TidUsed: v.tidUsed, + LastKeyIdx: v.lastKeyIdx, + KeyIdxUsed: v.keyIdxUsed, + } +} + +func (v *VP8Munger) SeedLast(state VP8MungerState) { + v.extLastPictureId = state.ExtLastPictureId + v.pictureIdUsed = state.PictureIdUsed + v.lastTl0PicIdx = state.LastTl0PicIdx + v.tl0PicIdxUsed = state.Tl0PicIdxUsed + v.tidUsed = state.TidUsed + v.lastKeyIdx = state.LastKeyIdx + v.keyIdxUsed = state.KeyIdxUsed +} + func (v *VP8Munger) SetLast(extPkt *buffer.ExtPacket) { vp8, ok := extPkt.Payload.(buffer.VP8) if !ok {