diff --git a/pkg/rtc/mediatracksubscriptions.go b/pkg/rtc/mediatracksubscriptions.go index c425695cd..a34271180 100644 --- a/pkg/rtc/mediatracksubscriptions.go +++ b/pkg/rtc/mediatracksubscriptions.go @@ -162,72 +162,110 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * } subTrack := NewSubscribedTrack(SubscribedTrackParams{ - PublisherID: t.params.MediaTrack.PublisherID(), - PublisherIdentity: t.params.MediaTrack.PublisherIdentity(), - SubscriberID: subscriberID, - SubscriberIdentity: sub.Identity(), - MediaTrack: t.params.MediaTrack, - DownTrack: downTrack, - AdaptiveStream: sub.GetAdaptiveStream(), + PublisherID: t.params.MediaTrack.PublisherID(), + PublisherIdentity: t.params.MediaTrack.PublisherIdentity(), + Subscriber: sub, + MediaTrack: t.params.MediaTrack, + DownTrack: downTrack, + AdaptiveStream: sub.GetAdaptiveStream(), + }) + + // Bind callback can happen from replaceTrack, so set it up early + downTrack.OnBind(func() { + wr.DetermineReceiver(downTrack.Codec()) + if err = wr.AddDownTrack(downTrack); err != nil { + t.params.Logger.Errorw("could not add down track", err, "participant", sub.Identity(), "pID", sub.ID()) + } + + go subTrack.Bound() + + // when down track is bound, start loop to send reports + go t.sendDownTrackBindingReports(sub) + + // initialize to default layer + t.notifySubscriberMaxQuality(subscriberID, downTrack.Codec(), livekit.VideoQuality_HIGH) + subTrack.SetPublisherMuted(t.params.MediaTrack.IsMuted()) }) var transceiver *webrtc.RTPTransceiver var sender *webrtc.RTPSender - if sub.ProtocolVersion().SupportsTransceiverReuse() { - // - // AddTrack will create a new transceiver or re-use an unused one - // if the attributes match. This prevents SDP from bloating - // because of dormant transceivers building up. - // - sender, err = sub.SubscriberPC().AddTrack(downTrack) - if err != nil { - return nil, err - } - // as there is no way to get transceiver from sender, search - for _, tr := range sub.SubscriberPC().GetTransceivers() { - if tr.Sender() == sender { - transceiver = tr - break + // try cached RTP senders for a chance to replace track + replacedTrack := false + existingTransceiver := sub.GetCachedRTPTransceiver(trackID) + if existingTransceiver != nil { + rtpSender := existingTransceiver.Sender() + if rtpSender != nil { + err := rtpSender.ReplaceTrack(downTrack) + if err == nil { + sender = rtpSender + transceiver = existingTransceiver + replacedTrack = true } } - if transceiver == nil { - // cannot add, no transceiver - return nil, errors.New("cannot subscribe without a transceiver in place") - } - } else { - transceiver, err = sub.SubscriberPC().AddTransceiverFromTrack(downTrack, webrtc.RTPTransceiverInit{ - Direction: webrtc.RTPTransceiverDirectionSendonly, - }) - if err != nil { - return nil, err - } - sender = transceiver.Sender() - if sender == nil { - // cannot add, no sender - return nil, errors.New("cannot subscribe without a sender in place") + if !replacedTrack { + // Could not re-use cached transceiver for this track. + // Stop the transceiver so that it is at least not active. + // It is not usable once stopped, + // + // Adding down track will create a new transceiver (or re-use + // an inactive existing one). In either case, a renegotiation + // will happen and that will notify remote of this stopped + // transceiver + existingTransceiver.Stop() } } + // if cannot replace, find an unused transceiver or add new one + if transceiver == nil { + if sub.ProtocolVersion().SupportsTransceiverReuse() { + // + // AddTrack will create a new transceiver or re-use an unused one + // if the attributes match. This prevents SDP from bloating + // because of dormant transceivers building up. + // + sender, err = sub.SubscriberPC().AddTrack(downTrack) + if err != nil { + return nil, err + } + + // as there is no way to get transceiver from sender, search + for _, tr := range sub.SubscriberPC().GetTransceivers() { + if tr.Sender() == sender { + transceiver = tr + break + } + } + } else { + transceiver, err = sub.SubscriberPC().AddTransceiverFromTrack(downTrack, webrtc.RTPTransceiverInit{ + Direction: webrtc.RTPTransceiverDirectionSendonly, + }) + if err != nil { + return nil, err + } + + sender = transceiver.Sender() + } + } + if transceiver == nil { + // cannot add, no transceiver + return nil, errors.New("cannot subscribe without a transceiver in place") + } + if sender == nil { + // cannot add, no sender + return nil, errors.New("cannot subscribe without a sender in place") + } + + // 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) + sendParameters := sender.GetParameters() downTrack.SetRTPHeaderExtensions(sendParameters.HeaderExtensions) downTrack.SetTransceiver(transceiver) - // when out track is bound, start loop to send reports - downTrack.OnBind(func() { - wr.DetermineReceiver(downTrack.Codec()) - if err = wr.AddDownTrack(downTrack); err != nil { - logger.Errorw("could not add down track", err, "participant", sub.Identity(), "pID", sub.ID()) - } - go subTrack.Bound() - go t.sendDownTrackBindingReports(sub) - // initialize to default layer - t.notifySubscriberMaxQuality(subscriberID, downTrack.Codec(), livekit.VideoQuality_HIGH) - subTrack.SetPublisherMuted(t.params.MediaTrack.IsMuted()) - }) - downTrack.OnStatsUpdate(func(_ *sfu.DownTrack, stat *livekit.AnalyticsStat) { t.params.Telemetry.TrackStats(livekit.StreamType_DOWNSTREAM, subscriberID, trackID, stat) }) @@ -251,7 +289,9 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * // since sub will lock, run it in a goroutine to avoid deadlocks go func() { sub.AddSubscribedTrack(subTrack) - sub.Negotiate(false) + if !replacedTrack { + sub.Negotiate(false) + } }() t.params.Telemetry.TrackSubscribed(context.Background(), subscriberID, t.params.MediaTrack.ToProto(), @@ -272,7 +312,7 @@ func (t *MediaTrackSubscriptions) RemoveSubscriber(participantID livekit.Partici t.subscribedTracksMu.Unlock() if subTrack != nil { - subTrack.DownTrack().CloseWithFlush(!willBeResumed) + t.closeSubscribedTrack(subTrack, willBeResumed) } } @@ -289,10 +329,27 @@ func (t *MediaTrackSubscriptions) RemoveAllSubscribers(willBeResumed bool) { t.subscribedTracksMu.Unlock() for _, subTrack := range subscribedTracks { - subTrack.DownTrack().CloseWithFlush(!willBeResumed) + t.closeSubscribedTrack(subTrack, willBeResumed) } } +func (t *MediaTrackSubscriptions) closeSubscribedTrack(subTrack types.SubscribedTrack, willBeResumed bool) { + dt := subTrack.DownTrack() + if dt == nil { + return + } + + if willBeResumed { + tr := dt.GetTransceiver() + if tr != nil { + sub := subTrack.Subscriber() + sub.CacheRTPTransceiver(subTrack.ID(), tr) + } + } + + dt.CloseWithFlush(!willBeResumed) +} + func (t *MediaTrackSubscriptions) ResyncAllSubscribers() { t.params.Logger.Debugw("resyncing all subscribers") diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 75416faf6..9c02eaea7 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -137,6 +137,8 @@ type ParticipantImpl struct { activeCounter atomic.Int32 firstConnected atomic.Bool iceConfig types.IceConfig + + cachedRTPTransceivers map[livekit.TrackID]*webrtc.RTPTransceiver } func NewParticipant(params ParticipantParams) (*ParticipantImpl, error) { @@ -159,6 +161,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), } p.version.Store(params.InitialVersion) p.migrateState.Store(types.MigrateStateInit) @@ -1908,3 +1911,30 @@ func (p *ParticipantImpl) setDowntracksConnected() { } } } + +func (p *ParticipantImpl) CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver) { + p.lock.Lock() + if existing := p.cachedRTPTransceivers[trackID]; existing != nil && existing != rtpTransceiver { + p.params.Logger.Infow("cached transceiver change", "trackID", trackID) + } + p.cachedRTPTransceivers[trackID] = rtpTransceiver + p.lock.Unlock() +} + +func (p *ParticipantImpl) UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver) { + p.lock.Lock() + for trackID, tr := range p.cachedRTPTransceivers { + if tr == rtpTransceiver { + delete(p.cachedRTPTransceivers, trackID) + break + } + } + p.lock.Unlock() +} + +func (p *ParticipantImpl) GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver { + p.lock.RLock() + defer p.lock.RUnlock() + + return p.cachedRTPTransceivers[trackID] +} diff --git a/pkg/rtc/subscribedtrack.go b/pkg/rtc/subscribedtrack.go index 7f9fc43d0..22a3b613d 100644 --- a/pkg/rtc/subscribedtrack.go +++ b/pkg/rtc/subscribedtrack.go @@ -19,13 +19,12 @@ const ( ) type SubscribedTrackParams struct { - PublisherID livekit.ParticipantID - PublisherIdentity livekit.ParticipantIdentity - SubscriberID livekit.ParticipantID - SubscriberIdentity livekit.ParticipantIdentity - MediaTrack types.MediaTrack - DownTrack *sfu.DownTrack - AdaptiveStream bool + PublisherID livekit.ParticipantID + PublisherIdentity livekit.ParticipantIdentity + Subscriber types.LocalParticipant + MediaTrack types.MediaTrack + DownTrack *sfu.DownTrack + AdaptiveStream bool } type SubscribedTrack struct { @@ -34,7 +33,8 @@ type SubscribedTrack struct { pubMuted atomic.Bool settings atomic.Value // *livekit.UpdateTrackSettings - onBind func() + onBind atomic.Value // func() + bound atomic.Bool debouncer func(func()) } @@ -49,15 +49,22 @@ func NewSubscribedTrack(params SubscribedTrackParams) *SubscribedTrack { } func (t *SubscribedTrack) OnBind(f func()) { - t.onBind = f + t.onBind.Store(f) + + t.maybeOnBind() } func (t *SubscribedTrack) Bound() { + t.bound.Store(true) if !t.params.AdaptiveStream { t.params.DownTrack.SetMaxSpatialLayer(utils.SpatialLayerForQuality(livekit.VideoQuality_HIGH)) } - if t.onBind != nil { - t.onBind() + t.maybeOnBind() +} + +func (t *SubscribedTrack) maybeOnBind() { + if onBind := t.onBind.Load(); onBind != nil && t.bound.Load() { + go onBind.(func())() } } @@ -74,11 +81,15 @@ func (t *SubscribedTrack) PublisherIdentity() livekit.ParticipantIdentity { } func (t *SubscribedTrack) SubscriberID() livekit.ParticipantID { - return t.params.SubscriberID + return t.params.Subscriber.ID() } func (t *SubscribedTrack) SubscriberIdentity() livekit.ParticipantIdentity { - return t.params.SubscriberIdentity + return t.params.Subscriber.Identity() +} + +func (t *SubscribedTrack) Subscriber() types.LocalParticipant { + return t.params.Subscriber } func (t *SubscribedTrack) DownTrack() *sfu.DownTrack { diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 0dc1895b6..a0710a81b 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -246,6 +246,10 @@ type LocalParticipant interface { SetMigrateInfo(previousAnswer *webrtc.SessionDescription, mediaTracks []*livekit.TrackPublishedResponse, dataChannels []*livekit.DataChannelInfo) UpdateRTT(rtt uint32) + + CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver) + UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver) + GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver } // Room is a container of participants, and can provide room-level actions @@ -325,6 +329,7 @@ type SubscribedTrack interface { PublisherIdentity() livekit.ParticipantIdentity SubscriberID() livekit.ParticipantID SubscriberIdentity() livekit.ParticipantIdentity + Subscriber() LocalParticipant DownTrack() *sfu.DownTrack MediaTrack() MediaTrack IsMuted() bool diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index 50993aa37..219936036 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -50,6 +50,12 @@ type FakeLocalParticipant struct { addTrackArgsForCall []struct { arg1 *livekit.AddTrackRequest } + CacheRTPTransceiverStub func(livekit.TrackID, *webrtc.RTPTransceiver) + cacheRTPTransceiverMutex sync.RWMutex + cacheRTPTransceiverArgsForCall []struct { + arg1 livekit.TrackID + arg2 *webrtc.RTPTransceiver + } CanPublishStub func() bool canPublishMutex sync.RWMutex canPublishArgsForCall []struct { @@ -144,6 +150,17 @@ type FakeLocalParticipant struct { result1 float64 result2 bool } + GetCachedRTPTransceiverStub func(livekit.TrackID) *webrtc.RTPTransceiver + getCachedRTPTransceiverMutex sync.RWMutex + getCachedRTPTransceiverArgsForCall []struct { + arg1 livekit.TrackID + } + getCachedRTPTransceiverReturns struct { + result1 *webrtc.RTPTransceiver + } + getCachedRTPTransceiverReturnsOnCall map[int]struct { + result1 *webrtc.RTPTransceiver + } GetConnectionQualityStub func() *livekit.ConnectionQualityInfo getConnectionQualityMutex sync.RWMutex getConnectionQualityArgsForCall []struct { @@ -589,6 +606,11 @@ type FakeLocalParticipant struct { toProtoReturnsOnCall map[int]struct { result1 *livekit.ParticipantInfo } + UncacheRTPTransceiverStub func(*webrtc.RTPTransceiver) + uncacheRTPTransceiverMutex sync.RWMutex + uncacheRTPTransceiverArgsForCall []struct { + arg1 *webrtc.RTPTransceiver + } UpdateMediaLossStub func(livekit.NodeID, livekit.TrackID, uint32) error updateMediaLossMutex sync.RWMutex updateMediaLossArgsForCall []struct { @@ -851,6 +873,39 @@ 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 { + arg1 livekit.TrackID + arg2 *webrtc.RTPTransceiver + }{arg1, arg2}) + stub := fake.CacheRTPTransceiverStub + fake.recordInvocation("CacheRTPTransceiver", []interface{}{arg1, arg2}) + fake.cacheRTPTransceiverMutex.Unlock() + if stub != nil { + fake.CacheRTPTransceiverStub(arg1, arg2) + } +} + +func (fake *FakeLocalParticipant) CacheRTPTransceiverCallCount() int { + fake.cacheRTPTransceiverMutex.RLock() + defer fake.cacheRTPTransceiverMutex.RUnlock() + return len(fake.cacheRTPTransceiverArgsForCall) +} + +func (fake *FakeLocalParticipant) CacheRTPTransceiverCalls(stub func(livekit.TrackID, *webrtc.RTPTransceiver)) { + fake.cacheRTPTransceiverMutex.Lock() + defer fake.cacheRTPTransceiverMutex.Unlock() + fake.CacheRTPTransceiverStub = 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) CanPublish() bool { fake.canPublishMutex.Lock() ret, specificReturn := fake.canPublishReturnsOnCall[len(fake.canPublishArgsForCall)] @@ -1340,6 +1395,67 @@ 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 { + arg1 livekit.TrackID + }{arg1}) + stub := fake.GetCachedRTPTransceiverStub + fakeReturns := fake.getCachedRTPTransceiverReturns + fake.recordInvocation("GetCachedRTPTransceiver", []interface{}{arg1}) + fake.getCachedRTPTransceiverMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCallCount() int { + fake.getCachedRTPTransceiverMutex.RLock() + defer fake.getCachedRTPTransceiverMutex.RUnlock() + return len(fake.getCachedRTPTransceiverArgsForCall) +} + +func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCalls(stub func(livekit.TrackID) *webrtc.RTPTransceiver) { + fake.getCachedRTPTransceiverMutex.Lock() + defer fake.getCachedRTPTransceiverMutex.Unlock() + fake.GetCachedRTPTransceiverStub = stub +} + +func (fake *FakeLocalParticipant) GetCachedRTPTransceiverArgsForCall(i int) livekit.TrackID { + fake.getCachedRTPTransceiverMutex.RLock() + defer fake.getCachedRTPTransceiverMutex.RUnlock() + argsForCall := fake.getCachedRTPTransceiverArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakeLocalParticipant) GetCachedRTPTransceiverReturns(result1 *webrtc.RTPTransceiver) { + fake.getCachedRTPTransceiverMutex.Lock() + defer fake.getCachedRTPTransceiverMutex.Unlock() + fake.GetCachedRTPTransceiverStub = nil + fake.getCachedRTPTransceiverReturns = struct { + result1 *webrtc.RTPTransceiver + }{result1} +} + +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 { + result1 *webrtc.RTPTransceiver + }) + } + fake.getCachedRTPTransceiverReturnsOnCall[i] = struct { + result1 *webrtc.RTPTransceiver + }{result1} +} + func (fake *FakeLocalParticipant) GetConnectionQuality() *livekit.ConnectionQualityInfo { fake.getConnectionQualityMutex.Lock() ret, specificReturn := fake.getConnectionQualityReturnsOnCall[len(fake.getConnectionQualityArgsForCall)] @@ -3805,6 +3921,38 @@ 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 { + arg1 *webrtc.RTPTransceiver + }{arg1}) + stub := fake.UncacheRTPTransceiverStub + fake.recordInvocation("UncacheRTPTransceiver", []interface{}{arg1}) + fake.uncacheRTPTransceiverMutex.Unlock() + if stub != nil { + fake.UncacheRTPTransceiverStub(arg1) + } +} + +func (fake *FakeLocalParticipant) UncacheRTPTransceiverCallCount() int { + fake.uncacheRTPTransceiverMutex.RLock() + defer fake.uncacheRTPTransceiverMutex.RUnlock() + return len(fake.uncacheRTPTransceiverArgsForCall) +} + +func (fake *FakeLocalParticipant) UncacheRTPTransceiverCalls(stub func(*webrtc.RTPTransceiver)) { + fake.uncacheRTPTransceiverMutex.Lock() + defer fake.uncacheRTPTransceiverMutex.Unlock() + fake.UncacheRTPTransceiverStub = stub +} + +func (fake *FakeLocalParticipant) UncacheRTPTransceiverArgsForCall(i int) *webrtc.RTPTransceiver { + fake.uncacheRTPTransceiverMutex.RLock() + defer fake.uncacheRTPTransceiverMutex.RUnlock() + argsForCall := fake.uncacheRTPTransceiverArgsForCall[i] + return argsForCall.arg1 +} + func (fake *FakeLocalParticipant) UpdateMediaLoss(arg1 livekit.NodeID, arg2 livekit.TrackID, arg3 uint32) error { fake.updateMediaLossMutex.Lock() ret, specificReturn := fake.updateMediaLossReturnsOnCall[len(fake.updateMediaLossArgsForCall)] @@ -4165,6 +4313,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.canPublishMutex.RLock() defer fake.canPublishMutex.RUnlock() fake.canPublishDataMutex.RLock() @@ -4183,6 +4333,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.getConnectionQualityMutex.RLock() defer fake.getConnectionQualityMutex.RUnlock() fake.getLoggerMutex.RLock() @@ -4285,6 +4437,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.updateMediaLossMutex.RLock() defer fake.updateMediaLossMutex.RUnlock() fake.updateRTTMutex.RLock() diff --git a/pkg/rtc/types/typesfakes/fake_subscribed_track.go b/pkg/rtc/types/typesfakes/fake_subscribed_track.go index f06924a72..2d0dc8d5e 100644 --- a/pkg/rtc/types/typesfakes/fake_subscribed_track.go +++ b/pkg/rtc/types/typesfakes/fake_subscribed_track.go @@ -80,6 +80,16 @@ type FakeSubscribedTrack struct { setPublisherMutedArgsForCall []struct { arg1 bool } + SubscriberStub func() types.LocalParticipant + subscriberMutex sync.RWMutex + subscriberArgsForCall []struct { + } + subscriberReturns struct { + result1 types.LocalParticipant + } + subscriberReturnsOnCall map[int]struct { + result1 types.LocalParticipant + } SubscriberIDStub func() livekit.ParticipantID subscriberIDMutex sync.RWMutex subscriberIDArgsForCall []struct { @@ -495,6 +505,59 @@ func (fake *FakeSubscribedTrack) SetPublisherMutedArgsForCall(i int) bool { return argsForCall.arg1 } +func (fake *FakeSubscribedTrack) Subscriber() types.LocalParticipant { + fake.subscriberMutex.Lock() + ret, specificReturn := fake.subscriberReturnsOnCall[len(fake.subscriberArgsForCall)] + fake.subscriberArgsForCall = append(fake.subscriberArgsForCall, struct { + }{}) + stub := fake.SubscriberStub + fakeReturns := fake.subscriberReturns + fake.recordInvocation("Subscriber", []interface{}{}) + fake.subscriberMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeSubscribedTrack) SubscriberCallCount() int { + fake.subscriberMutex.RLock() + defer fake.subscriberMutex.RUnlock() + return len(fake.subscriberArgsForCall) +} + +func (fake *FakeSubscribedTrack) SubscriberCalls(stub func() types.LocalParticipant) { + fake.subscriberMutex.Lock() + defer fake.subscriberMutex.Unlock() + fake.SubscriberStub = stub +} + +func (fake *FakeSubscribedTrack) SubscriberReturns(result1 types.LocalParticipant) { + fake.subscriberMutex.Lock() + defer fake.subscriberMutex.Unlock() + fake.SubscriberStub = nil + fake.subscriberReturns = struct { + result1 types.LocalParticipant + }{result1} +} + +func (fake *FakeSubscribedTrack) SubscriberReturnsOnCall(i int, result1 types.LocalParticipant) { + fake.subscriberMutex.Lock() + defer fake.subscriberMutex.Unlock() + fake.SubscriberStub = nil + if fake.subscriberReturnsOnCall == nil { + fake.subscriberReturnsOnCall = make(map[int]struct { + result1 types.LocalParticipant + }) + } + fake.subscriberReturnsOnCall[i] = struct { + result1 types.LocalParticipant + }{result1} +} + func (fake *FakeSubscribedTrack) SubscriberID() livekit.ParticipantID { fake.subscriberIDMutex.Lock() ret, specificReturn := fake.subscriberIDReturnsOnCall[len(fake.subscriberIDArgsForCall)] @@ -676,6 +739,8 @@ func (fake *FakeSubscribedTrack) Invocations() map[string][][]interface{} { defer fake.publisherIdentityMutex.RUnlock() fake.setPublisherMutedMutex.RLock() defer fake.setPublisherMutedMutex.RUnlock() + fake.subscriberMutex.RLock() + defer fake.subscriberMutex.RUnlock() fake.subscriberIDMutex.RLock() defer fake.subscriberIDMutex.RUnlock() fake.subscriberIdentityMutex.RLock() diff --git a/pkg/rtc/wrappedreceiver.go b/pkg/rtc/wrappedreceiver.go index 7cf5ec70e..6a42600b3 100644 --- a/pkg/rtc/wrappedreceiver.go +++ b/pkg/rtc/wrappedreceiver.go @@ -67,8 +67,15 @@ type DummyReceiver struct { streamId string codec webrtc.RTPCodecParameters headerExtensions []webrtc.RTPHeaderExtensionParameter - downtrackLock sync.Mutex - downtracks map[livekit.ParticipantID]sfu.TrackSender + + downtrackLock sync.Mutex + downtracks map[livekit.ParticipantID]sfu.TrackSender + + settingsLock sync.Mutex + maxExpectedLayerValid bool + maxExpectedLayer int32 + pausedValid bool + paused bool } func NewDummyReceiver(trackID livekit.TrackID, streamId string, codec webrtc.RTPCodecParameters, headerExtensions []webrtc.RTPHeaderExtensionParameter) *DummyReceiver { @@ -87,13 +94,25 @@ func (d *DummyReceiver) Receiver() sfu.TrackReceiver { } func (d *DummyReceiver) Upgrade(receiver sfu.TrackReceiver) { - d.downtrackLock.Lock() - defer d.downtrackLock.Unlock() d.receiver.CompareAndSwap(nil, receiver) + + d.downtrackLock.Lock() for _, t := range d.downtracks { receiver.AddDownTrack(t) } d.downtracks = make(map[livekit.ParticipantID]sfu.TrackSender) + d.downtrackLock.Unlock() + + d.settingsLock.Lock() + if d.maxExpectedLayerValid { + receiver.SetMaxExpectedSpatialLayer(d.maxExpectedLayer) + } + d.maxExpectedLayerValid = false + if d.pausedValid { + receiver.SetUpTrackPaused(d.paused) + } + d.pausedValid = false + d.settingsLock.Unlock() } func (d *DummyReceiver) TrackID() livekit.TrackID { @@ -148,12 +167,22 @@ func (d *DummyReceiver) SendPLI(layer int32, force bool) { func (d *DummyReceiver) SetUpTrackPaused(paused bool) { if r, ok := d.receiver.Load().(sfu.TrackReceiver); ok { r.SetUpTrackPaused(paused) + } else { + d.settingsLock.Lock() + d.pausedValid = true + d.paused = paused + d.settingsLock.Unlock() } } func (d *DummyReceiver) SetMaxExpectedSpatialLayer(layer int32) { if r, ok := d.receiver.Load().(sfu.TrackReceiver); ok { r.SetMaxExpectedSpatialLayer(layer) + } else { + d.settingsLock.Lock() + d.maxExpectedLayerValid = true + d.maxExpectedLayer = layer + d.settingsLock.Unlock() } } diff --git a/pkg/sfu/downtrack.go b/pkg/sfu/downtrack.go index 0a488319a..eba45dacc 100644 --- a/pkg/sfu/downtrack.go +++ b/pkg/sfu/downtrack.go @@ -364,6 +364,10 @@ func (d *DownTrack) SetTransceiver(transceiver *webrtc.RTPTransceiver) { d.transceiver = transceiver } +func (d *DownTrack) GetTransceiver() *webrtc.RTPTransceiver { + return d.transceiver +} + func (d *DownTrack) maybeStartKeyFrameRequester() { // // Always move to next generation to abandon any running key frame requester diff --git a/pkg/sfu/streamallocator.go b/pkg/sfu/streamallocator.go index b94d3c050..cbf01c263 100644 --- a/pkg/sfu/streamallocator.go +++ b/pkg/sfu/streamallocator.go @@ -240,7 +240,9 @@ func (s *StreamAllocator) AddTrack(downTrack *DownTrack, params AddTrackParams) func (s *StreamAllocator) RemoveTrack(downTrack *DownTrack) { s.videoTracksMu.Lock() - delete(s.videoTracks, livekit.TrackID(downTrack.ID())) + if existing := s.videoTracks[livekit.TrackID(downTrack.ID())]; existing != nil && existing.DownTrack() == downTrack { + delete(s.videoTracks, livekit.TrackID(downTrack.ID())) + } s.videoTracksMu.Unlock() // LK-TODO: use any saved bandwidth to re-distribute