diff --git a/pkg/rtc/subscriptionmanager.go b/pkg/rtc/subscriptionmanager.go index 38d939f22..b5edddca1 100644 --- a/pkg/rtc/subscriptionmanager.go +++ b/pkg/rtc/subscriptionmanager.go @@ -189,19 +189,19 @@ func (m *SubscriptionManager) SubscribeToTrack(trackID livekit.TrackID, isSync b return } - sub, desireChanged := m.setDesired(trackID, true) - if sub == nil { + // find or create and set desired under one lock, so that a concurrent subscribe, + // settings update or cleanup cannot replace or remove the subscription in between + m.lock.Lock() + sub, ok := m.subscriptions[trackID] + if !ok { sLogger := m.params.Logger.WithValues( "trackID", trackID, ) sub = newMediaTrackSubscription(m.params.Participant.ID(), trackID, sLogger) - - m.lock.Lock() m.subscriptions[trackID] = sub - m.lock.Unlock() - - sub, desireChanged = m.setDesired(trackID, true) } + desireChanged := sub.setDesired(true) + m.lock.Unlock() if desireChanged { sub.logger.Debugw("subscribing to track") } @@ -234,19 +234,18 @@ func (m *SubscriptionManager) SubscribeToDataTrack(trackID livekit.TrackID) { return } - sub, desireChanged := m.setDataTrackDesired(trackID, true) - if sub == nil { + // same as SubscribeToTrack, find or create and set desired under one lock + m.lock.Lock() + sub, ok := m.dataTrackSubscriptions[trackID] + if !ok { sLogger := m.params.Logger.WithValues( "trackID", trackID, ) sub = newDataTrackSubscription(m.params.Participant.ID(), trackID, sLogger) - - m.lock.Lock() m.dataTrackSubscriptions[trackID] = sub - m.lock.Unlock() - - sub, desireChanged = m.setDataTrackDesired(trackID, true) } + desireChanged := sub.setDesired(true) + m.lock.Unlock() if desireChanged { sub.logger.Debugw("subscribing to data track") } @@ -389,6 +388,7 @@ func (m *SubscriptionManager) UpdateSubscribedTrackSettings(trackID livekit.Trac sub = newMediaTrackSubscription(m.params.Participant.ID(), trackID, sLogger) m.subscriptions[trackID] = sub } + sub.keepForSubscribe() m.lock.Unlock() sub.setSettings(settings) @@ -404,6 +404,7 @@ func (m *SubscriptionManager) UpdateDataTrackSubscriptionOptions(trackID livekit sub = newDataTrackSubscription(m.params.Participant.ID(), trackID, sLogger) m.dataTrackSubscriptions[trackID] = sub } + sub.keepForSubscribe() m.lock.Unlock() sub.setSubscriptionOptions(subscriptionOptions) @@ -601,7 +602,7 @@ func (m *SubscriptionManager) reconcileSubscription(s *mediaTrackSubscription) { } m.lock.Lock() - if s.needsCleanup() { + if m.subscriptions[s.trackID] == s && s.needsCleanup() { s.logger.Debugw("cleanup removing subscription") delete(m.subscriptions, s.trackID) } @@ -673,15 +674,22 @@ func (m *SubscriptionManager) reconcileDataTrackSubscription(s *dataTrackSubscri s.logger.Warnw("failed to unsubscribe", err) } + // a subscribe during the removal sets the entry as desired again, keep it then m.lock.Lock() - delete(m.dataTrackSubscriptions, s.trackID) + removed := m.dataTrackSubscriptions[s.trackID] == s && !s.isDesired() + if removed { + delete(m.dataTrackSubscriptions, s.trackID) + } m.lock.Unlock() m.notifyDataTrackSubscriberHandles() + if !removed { + m.queueReconcileDataTrack(s.trackID) + } return } m.lock.Lock() - cleanedUp := s.needsCleanup() + cleanedUp := m.dataTrackSubscriptions[s.trackID] == s && s.needsCleanup() if cleanedUp { s.logger.Debugw("cleanup removing data track subscription") delete(m.dataTrackSubscriptions, s.trackID) @@ -1280,6 +1288,10 @@ type trackSubscription struct { // the timestamp when the subscription was started, will be reset when downtrack is closed with expected resume subscribeAt atomic.Pointer[time.Time] + + // an entry that is not desired is not cleaned up before this time, + // so that settings sent before a subscribe are there when the subscribe comes + keepUntil time.Time } // set permission and return true if it has changed @@ -1331,6 +1343,7 @@ func (s *trackSubscription) setDesired(desired bool) bool { t := time.Now() s.subStartedAt.Store(&t) s.subscribeAt.Store(&t) + s.keepUntil = time.Time{} } if s.desired == desired { @@ -1372,6 +1385,18 @@ func (s *trackSubscription) getNumAttempts() int32 { return s.numAttempts.Load() } +func (s *trackSubscription) keepForSubscribe() { + s.lock.Lock() + defer s.lock.Unlock() + if !s.desired { + s.keepUntil = time.Now().Add(notFoundTimeout) + } +} + +func (s *trackSubscription) isKeptLocked() bool { + return time.Now().Before(s.keepUntil) +} + func (s *trackSubscription) durationSinceStart() time.Duration { t := s.subStartedAt.Load() if t == nil { @@ -1604,7 +1629,7 @@ func (s *mediaTrackSubscription) needsBind() bool { func (s *mediaTrackSubscription) needsCleanup() bool { s.lock.RLock() defer s.lock.RUnlock() - return !s.desired && s.subscribedTrack == nil + return !s.desired && s.subscribedTrack == nil && !s.isKeptLocked() } // ----------------------------------------------------------------- @@ -1645,7 +1670,7 @@ func (s *dataTrackSubscription) needsUnsubscribe() bool { func (s *dataTrackSubscription) needsCleanup() bool { s.lock.RLock() defer s.lock.RUnlock() - return !s.desired && s.dataDownTrack == nil + return !s.desired && s.dataDownTrack == nil && !s.isKeptLocked() } func (s *dataTrackSubscription) setDataDownTrack(dataDownTrack types.DataDownTrack) { diff --git a/pkg/rtc/subscriptionmanager_test.go b/pkg/rtc/subscriptionmanager_test.go index 15ecc1096..b5305e32d 100644 --- a/pkg/rtc/subscriptionmanager_test.go +++ b/pkg/rtc/subscriptionmanager_test.go @@ -353,6 +353,141 @@ func TestUpdateSettingsBeforeSubscription(t *testing.T) { require.Equal(t, settings.Height, applied.Height) } +func TestConcurrentFirstSubscribe(t *testing.T) { + settings := &livekit.UpdateTrackSettings{Width: 100, Height: 100} + for range 200 { + sm := newTestSubscriptionManager() + + var lock sync.Mutex + var subscribed bool + mt := &typesfakes.FakeMediaTrack{} + mt.IDReturns("track") + mt.AddSubscriberCalls(func(types.LocalParticipant) (types.SubscribedTrack, error) { + lock.Lock() + defer lock.Unlock() + if subscribed { + return nil, errAlreadySubscribed + } + subscribed = true + st := &typesfakes.FakeSubscribedTrack{} + st.IDReturns("track") + st.MediaTrackReturns(mt) + return st, nil + }) + sm.params.TrackResolver = func(types.LocalParticipant, livekit.TrackID) types.MediaResolverResult { + return types.MediaResolverResult{ + Track: mt, + HasPermission: true, + PublisherID: "pubID", + PublisherIdentity: "pub", + TrackChangedNotifier: utils.NewChangeNotifier(), + TrackRemovedNotifier: utils.NewChangeNotifier(), + } + } + + // two first subscribes and a settings update race to create the subscription + var start, done sync.WaitGroup + start.Add(1) + for _, f := range []func(){ + func() { sm.SubscribeToTrack("track", false) }, + func() { sm.SubscribeToTrack("track", false) }, + func() { sm.UpdateSubscribedTrackSettings("track", settings) }, + } { + done.Add(1) + go func() { + defer done.Done() + start.Wait() + f() + }() + } + start.Done() + done.Wait() + + sm.lock.RLock() + s := sm.subscriptions["track"] + sm.lock.RUnlock() + require.Eventually(t, func() bool { + return !s.needsSubscribe() + }, subSettleTimeout, subCheckInterval, "the subscription that is kept should own the down track") + st := s.getSubscribedTrack().(*typesfakes.FakeSubscribedTrack) + require.Eventually(t, func() bool { + n := st.UpdateSubscriberSettingsCallCount() + if n == 0 { + return false + } + applied, _ := st.UpdateSubscriberSettingsArgsForCall(n - 1) + return applied == settings + }, subSettleTimeout, subCheckInterval, "the down track should get the settings") + + sm.Close(false) + } +} + +func TestSettingsKeptForSubscribe(t *testing.T) { + t.Run("media", func(t *testing.T) { + sm := newTestSubscriptionManager() + defer sm.Close(false) + resolver := newTestResolver(true, true, "pub", "pubID") + sm.params.TrackResolver = resolver.Resolve + + settings := &livekit.UpdateTrackSettings{Disabled: true} + sm.UpdateSubscribedTrackSettings("track", settings) + // a cleanup pass between the settings and the subscribe must not drop the settings + sm.reconcileSubscriptions() + sm.SubscribeToTrack("track", false) + + sm.lock.RLock() + s := sm.subscriptions["track"] + sm.lock.RUnlock() + require.Eventually(t, func() bool { + return !s.needsSubscribe() + }, subSettleTimeout, subCheckInterval, "track should be subscribed") + st := s.getSubscribedTrack().(*typesfakes.FakeSubscribedTrack) + require.Eventually(t, func() bool { + n := st.UpdateSubscriberSettingsCallCount() + if n == 0 { + return false + } + applied, _ := st.UpdateSubscriberSettingsArgsForCall(n - 1) + return applied == settings + }, subSettleTimeout, subCheckInterval, "the down track should get the settings") + + // settings for a track that is never subscribed are cleaned up after notFoundTimeout + sm.UpdateSubscribedTrackSettings("other", settings) + require.Eventually(t, func() bool { + sm.lock.RLock() + defer sm.lock.RUnlock() + _, ok := sm.subscriptions["other"] + return !ok + }, subSettleTimeout, subCheckInterval, "unused settings should be cleaned up") + }) + + t.Run("data", func(t *testing.T) { + sm := newTestSubscriptionManager() + defer sm.Close(false) + resolver := newTestDataTrackResolver(true, true, "pub", "pubID") + sm.params.DataTrackResolver = resolver.Resolve + + fps := uint32(5) + options := &livekit.DataTrackSubscriptionOptions{TargetFps: &fps} + sm.UpdateDataTrackSubscriptionOptions("track", options) + // a cleanup pass between the options and the subscribe must not drop the options + sm.reconcileDataTrackSubscriptions() + sm.SubscribeToDataTrack("track") + + sm.lock.RLock() + s := sm.dataTrackSubscriptions["track"] + sm.lock.RUnlock() + require.Eventually(t, func() bool { + return s.getDataDownTrack() != nil + }, subSettleTimeout, subCheckInterval, "data track should be subscribed") + ddt := s.getDataDownTrack().(*typesfakes.FakeDataDownTrack) + n := ddt.UpdateSubscriptionOptionsCallCount() + require.NotZero(t, n) + require.Equal(t, options, ddt.UpdateSubscriptionOptionsArgsForCall(n-1)) + }) +} + func TestSubscriptionLimits(t *testing.T) { sm := newTestSubscriptionManagerWithParams(testSubscriptionParams{ SubscriptionLimitAudio: 1, @@ -527,6 +662,37 @@ func TestSubscribeDataTrack(t *testing.T) { require.Equal(t, 2, resolver.dataTrack.AddSubscriberCallCount()) }) + t.Run("subscribe again during unsubscribe", func(t *testing.T) { + sm := newTestSubscriptionManager() + defer sm.Close(false) + resolver := newTestDataTrackResolver(true, true, "pub", "pubID") + sm.params.DataTrackResolver = resolver.Resolve + + sm.SubscribeToDataTrack("track") + sm.lock.RLock() + s := sm.dataTrackSubscriptions["track"] + sm.lock.RUnlock() + require.Eventually(t, func() bool { + return !s.needsSubscribe() + }, subSettleTimeout, subCheckInterval, "should be subscribed") + + // the client subscribes again while the removal runs, + // then the down track closes, as DataTrack.RemoveSubscriber does + ddt := s.getDataDownTrack().(*typesfakes.FakeDataDownTrack) + resolver.dataTrack.RemoveSubscriberCalls(func(livekit.ParticipantID) { + sm.SubscribeToDataTrack("track") + ddt.OnCloseArgsForCall(0)() + }) + sm.UnsubscribeFromDataTrack("track") + + require.Eventually(t, func() bool { + return resolver.dataTrack.AddSubscriberCallCount() == 2 && s.getDataDownTrack() != nil + }, subSettleTimeout, subCheckInterval, "should be subscribed again") + sm.lock.RLock() + require.Same(t, s, sm.dataTrackSubscriptions["track"]) + sm.lock.RUnlock() + }) + t.Run("unsubscribe before data track resolves", func(t *testing.T) { sm := newTestSubscriptionManager() defer sm.Close(false)