diff --git a/pkg/rtc/mediatrack.go b/pkg/rtc/mediatrack.go index 1def3a6bd..e7a248d01 100644 --- a/pkg/rtc/mediatrack.go +++ b/pkg/rtc/mediatrack.go @@ -242,6 +242,7 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track *webrtc.Tra ) newWR.SetRTCPCh(t.params.RTCPChan) newWR.OnCloseHandler(func() { + t.MediaTrackReceiver.SetClosing() t.MediaTrackReceiver.ClearReceiver(mime, false) if t.MediaTrackReceiver.TryClose() { if t.dynacastManager != nil { @@ -349,6 +350,7 @@ func (t *MediaTrack) Restart() { } func (t *MediaTrack) Close(willBeResumed bool) { + t.MediaTrackReceiver.SetClosing() if t.dynacastManager != nil { t.dynacastManager.Close() } diff --git a/pkg/rtc/mediatrackreceiver.go b/pkg/rtc/mediatrackreceiver.go index 625f1805d..10d2eeeab 100644 --- a/pkg/rtc/mediatrackreceiver.go +++ b/pkg/rtc/mediatrackreceiver.go @@ -36,6 +36,7 @@ type mediaTrackReceiverState int const ( mediaTrackReceiverStateOpen mediaTrackReceiverState = iota + mediaTrackReceiverStateClosing mediaTrackReceiverStateClosed ) @@ -43,6 +44,8 @@ func (m mediaTrackReceiverState) String() string { switch m { case mediaTrackReceiverStateOpen: return "OPEN" + case mediaTrackReceiverStateClosing: + return "CLOSING" case mediaTrackReceiverStateClosed: return "CLOSED" default: @@ -291,6 +294,20 @@ func (t *MediaTrackReceiver) OnVideoLayerUpdate(f func(layers []*livekit.VideoLa t.onVideoLayerUpdate = f } +func (t *MediaTrackReceiver) IsOpen() bool { + t.lock.RLock() + defer t.lock.RUnlock() + return t.state == mediaTrackReceiverStateOpen +} + +func (t *MediaTrackReceiver) SetClosing() { + t.lock.Lock() + defer t.lock.Unlock() + if t.state == mediaTrackReceiverStateOpen { + t.state = mediaTrackReceiverStateClosing + } +} + func (t *MediaTrackReceiver) TryClose() bool { t.lock.RLock() if t.state == mediaTrackReceiverStateClosed { diff --git a/pkg/rtc/roomtrackmanager.go b/pkg/rtc/roomtrackmanager.go index 1f119dbb5..a2813a04d 100644 --- a/pkg/rtc/roomtrackmanager.go +++ b/pkg/rtc/roomtrackmanager.go @@ -82,7 +82,15 @@ func (r *RoomTrackManager) GetTrackInfo(trackID livekit.TrackID) *TrackInfo { r.lock.RLock() defer r.lock.RUnlock() - return r.tracks[trackID] + info := r.tracks[trackID] + if info == nil { + return nil + } + // when track is about to close, do not resolve + if info.Track != nil && !info.Track.IsOpen() { + return nil + } + return info } func (r *RoomTrackManager) NotifyTrackChanged(trackID livekit.TrackID) { diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index a2251bc6a..4eef964d7 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -372,6 +372,7 @@ type MediaTrack interface { IsSimulcast() bool Close(willBeResumed bool) + IsOpen() bool // callbacks AddOnClose(func()) diff --git a/pkg/rtc/types/typesfakes/fake_local_media_track.go b/pkg/rtc/types/typesfakes/fake_local_media_track.go index 117c02d10..060712282 100644 --- a/pkg/rtc/types/typesfakes/fake_local_media_track.go +++ b/pkg/rtc/types/typesfakes/fake_local_media_track.go @@ -136,6 +136,16 @@ type FakeLocalMediaTrack struct { isMutedReturnsOnCall map[int]struct { result1 bool } + IsOpenStub func() bool + isOpenMutex sync.RWMutex + isOpenArgsForCall []struct { + } + isOpenReturns struct { + result1 bool + } + isOpenReturnsOnCall map[int]struct { + result1 bool + } IsSimulcastStub func() bool isSimulcastMutex sync.RWMutex isSimulcastArgsForCall []struct { @@ -966,6 +976,59 @@ func (fake *FakeLocalMediaTrack) IsMutedReturnsOnCall(i int, result1 bool) { }{result1} } +func (fake *FakeLocalMediaTrack) IsOpen() bool { + fake.isOpenMutex.Lock() + ret, specificReturn := fake.isOpenReturnsOnCall[len(fake.isOpenArgsForCall)] + fake.isOpenArgsForCall = append(fake.isOpenArgsForCall, struct { + }{}) + stub := fake.IsOpenStub + fakeReturns := fake.isOpenReturns + fake.recordInvocation("IsOpen", []interface{}{}) + fake.isOpenMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeLocalMediaTrack) IsOpenCallCount() int { + fake.isOpenMutex.RLock() + defer fake.isOpenMutex.RUnlock() + return len(fake.isOpenArgsForCall) +} + +func (fake *FakeLocalMediaTrack) IsOpenCalls(stub func() bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = stub +} + +func (fake *FakeLocalMediaTrack) IsOpenReturns(result1 bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = nil + fake.isOpenReturns = struct { + result1 bool + }{result1} +} + +func (fake *FakeLocalMediaTrack) IsOpenReturnsOnCall(i int, result1 bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = nil + if fake.isOpenReturnsOnCall == nil { + fake.isOpenReturnsOnCall = make(map[int]struct { + result1 bool + }) + } + fake.isOpenReturnsOnCall[i] = struct { + result1 bool + }{result1} +} + func (fake *FakeLocalMediaTrack) IsSimulcast() bool { fake.isSimulcastMutex.Lock() ret, specificReturn := fake.isSimulcastReturnsOnCall[len(fake.isSimulcastArgsForCall)] @@ -1881,6 +1944,8 @@ func (fake *FakeLocalMediaTrack) Invocations() map[string][][]interface{} { defer fake.iDMutex.RUnlock() fake.isMutedMutex.RLock() defer fake.isMutedMutex.RUnlock() + fake.isOpenMutex.RLock() + defer fake.isOpenMutex.RUnlock() fake.isSimulcastMutex.RLock() defer fake.isSimulcastMutex.RUnlock() fake.isSubscriberMutex.RLock() diff --git a/pkg/rtc/types/typesfakes/fake_media_track.go b/pkg/rtc/types/typesfakes/fake_media_track.go index e77fa4e38..de0e870f4 100644 --- a/pkg/rtc/types/typesfakes/fake_media_track.go +++ b/pkg/rtc/types/typesfakes/fake_media_track.go @@ -103,6 +103,16 @@ type FakeMediaTrack struct { isMutedReturnsOnCall map[int]struct { result1 bool } + IsOpenStub func() bool + isOpenMutex sync.RWMutex + isOpenArgsForCall []struct { + } + isOpenReturns struct { + result1 bool + } + isOpenReturnsOnCall map[int]struct { + result1 bool + } IsSimulcastStub func() bool isSimulcastMutex sync.RWMutex isSimulcastArgsForCall []struct { @@ -732,6 +742,59 @@ func (fake *FakeMediaTrack) IsMutedReturnsOnCall(i int, result1 bool) { }{result1} } +func (fake *FakeMediaTrack) IsOpen() bool { + fake.isOpenMutex.Lock() + ret, specificReturn := fake.isOpenReturnsOnCall[len(fake.isOpenArgsForCall)] + fake.isOpenArgsForCall = append(fake.isOpenArgsForCall, struct { + }{}) + stub := fake.IsOpenStub + fakeReturns := fake.isOpenReturns + fake.recordInvocation("IsOpen", []interface{}{}) + fake.isOpenMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeMediaTrack) IsOpenCallCount() int { + fake.isOpenMutex.RLock() + defer fake.isOpenMutex.RUnlock() + return len(fake.isOpenArgsForCall) +} + +func (fake *FakeMediaTrack) IsOpenCalls(stub func() bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = stub +} + +func (fake *FakeMediaTrack) IsOpenReturns(result1 bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = nil + fake.isOpenReturns = struct { + result1 bool + }{result1} +} + +func (fake *FakeMediaTrack) IsOpenReturnsOnCall(i int, result1 bool) { + fake.isOpenMutex.Lock() + defer fake.isOpenMutex.Unlock() + fake.IsOpenStub = nil + if fake.isOpenReturnsOnCall == nil { + fake.isOpenReturnsOnCall = make(map[int]struct { + result1 bool + }) + } + fake.isOpenReturnsOnCall[i] = struct { + result1 bool + }{result1} +} + func (fake *FakeMediaTrack) IsSimulcast() bool { fake.isSimulcastMutex.Lock() ret, specificReturn := fake.isSimulcastReturnsOnCall[len(fake.isSimulcastArgsForCall)] @@ -1461,6 +1524,8 @@ func (fake *FakeMediaTrack) Invocations() map[string][][]interface{} { defer fake.iDMutex.RUnlock() fake.isMutedMutex.RLock() defer fake.isMutedMutex.RUnlock() + fake.isOpenMutex.RLock() + defer fake.isOpenMutex.RUnlock() fake.isSimulcastMutex.RLock() defer fake.isSimulcastMutex.RUnlock() fake.isSubscriberMutex.RLock()