From 019ad88b08c2038498763e0bb55e5d417474145e Mon Sep 17 00:00:00 2001 From: Raja Subramanian Date: Sun, 17 Sep 2023 14:00:09 +0530 Subject: [PATCH] Do not force reconnect on resume if there is a pending track (#2081) * Do not force reconnect on resume if there is a pending track * move GetPendingTrack -> LocalParticipant --- pkg/rtc/participant.go | 13 ++++ pkg/rtc/room.go | 4 + pkg/rtc/types/interfaces.go | 3 +- .../typesfakes/fake_local_participant.go | 74 +++++++++++++++++++ 4 files changed, 93 insertions(+), 1 deletion(-) diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 39cce9073..07ae3036b 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -1637,6 +1637,19 @@ func (p *ParticipantImpl) addPendingTrackLocked(req *livekit.AddTrackRequest) *l return ti } +func (p *ParticipantImpl) GetPendingTrack(trackID livekit.TrackID) *livekit.TrackInfo { + p.pendingTracksLock.RLock() + defer p.pendingTracksLock.RUnlock() + + for _, t := range p.pendingTracks { + if livekit.TrackID(t.trackInfos[0].Sid) == trackID { + return t.trackInfos[0] + } + } + + return nil +} + func (p *ParticipantImpl) sendTrackPublished(cid string, ti *livekit.TrackInfo) { p.pubLogger.Debugw("sending track published", "cid", cid, "trackInfo", ti.String()) _ = p.writeMessage(&livekit.SignalResponse{ diff --git a/pkg/rtc/room.go b/pkg/rtc/room.go index 1589ada96..92f50ca01 100644 --- a/pkg/rtc/room.go +++ b/pkg/rtc/room.go @@ -581,6 +581,10 @@ func (r *Room) SyncState(participant types.LocalParticipant, state *livekit.Sync break } } + if !found { + // is there a pending track? + found = participant.GetPendingTrack(livekit.TrackID(ti.Sid)) != nil + } if !found { pLogger.Warnw("unknown track during resume", nil, "trackID", ti.Sid) shouldReconnect = true diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index e5e44d992..0b0e003ad 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -245,7 +245,7 @@ type Participant interface { SetMetadata(metadata string) IsPublisher() bool - GetPublishedTrack(sid livekit.TrackID) MediaTrack + GetPublishedTrack(trackID livekit.TrackID) MediaTrack GetPublishedTracks() []MediaTrack RemovePublishedTrack(track MediaTrack, willBeResumed bool, shouldClose bool) @@ -315,6 +315,7 @@ type LocalParticipant interface { GetICEConnectionType() ICEConnectionType GetBufferFactory() *buffer.Factory GetPlayoutDelayConfig() *livekit.PlayoutDelay + GetPendingTrack(trackID livekit.TrackID) *livekit.TrackInfo SetResponseSink(sink routing.MessageSink) CloseSignalConnection(reason SignallingCloseReason) diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index b5338c33a..fedebacf5 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -263,6 +263,17 @@ type FakeLocalParticipant struct { getPacerReturnsOnCall map[int]struct { result1 pacer.Pacer } + GetPendingTrackStub func(livekit.TrackID) *livekit.TrackInfo + getPendingTrackMutex sync.RWMutex + getPendingTrackArgsForCall []struct { + arg1 livekit.TrackID + } + getPendingTrackReturns struct { + result1 *livekit.TrackInfo + } + getPendingTrackReturnsOnCall map[int]struct { + result1 *livekit.TrackInfo + } GetPlayoutDelayConfigStub func() *livekit.PlayoutDelay getPlayoutDelayConfigMutex sync.RWMutex getPlayoutDelayConfigArgsForCall []struct { @@ -2169,6 +2180,67 @@ func (fake *FakeLocalParticipant) GetPacerReturnsOnCall(i int, result1 pacer.Pac }{result1} } +func (fake *FakeLocalParticipant) GetPendingTrack(arg1 livekit.TrackID) *livekit.TrackInfo { + fake.getPendingTrackMutex.Lock() + ret, specificReturn := fake.getPendingTrackReturnsOnCall[len(fake.getPendingTrackArgsForCall)] + fake.getPendingTrackArgsForCall = append(fake.getPendingTrackArgsForCall, struct { + arg1 livekit.TrackID + }{arg1}) + stub := fake.GetPendingTrackStub + fakeReturns := fake.getPendingTrackReturns + fake.recordInvocation("GetPendingTrack", []interface{}{arg1}) + fake.getPendingTrackMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeLocalParticipant) GetPendingTrackCallCount() int { + fake.getPendingTrackMutex.RLock() + defer fake.getPendingTrackMutex.RUnlock() + return len(fake.getPendingTrackArgsForCall) +} + +func (fake *FakeLocalParticipant) GetPendingTrackCalls(stub func(livekit.TrackID) *livekit.TrackInfo) { + fake.getPendingTrackMutex.Lock() + defer fake.getPendingTrackMutex.Unlock() + fake.GetPendingTrackStub = stub +} + +func (fake *FakeLocalParticipant) GetPendingTrackArgsForCall(i int) livekit.TrackID { + fake.getPendingTrackMutex.RLock() + defer fake.getPendingTrackMutex.RUnlock() + argsForCall := fake.getPendingTrackArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakeLocalParticipant) GetPendingTrackReturns(result1 *livekit.TrackInfo) { + fake.getPendingTrackMutex.Lock() + defer fake.getPendingTrackMutex.Unlock() + fake.GetPendingTrackStub = nil + fake.getPendingTrackReturns = struct { + result1 *livekit.TrackInfo + }{result1} +} + +func (fake *FakeLocalParticipant) GetPendingTrackReturnsOnCall(i int, result1 *livekit.TrackInfo) { + fake.getPendingTrackMutex.Lock() + defer fake.getPendingTrackMutex.Unlock() + fake.GetPendingTrackStub = nil + if fake.getPendingTrackReturnsOnCall == nil { + fake.getPendingTrackReturnsOnCall = make(map[int]struct { + result1 *livekit.TrackInfo + }) + } + fake.getPendingTrackReturnsOnCall[i] = struct { + result1 *livekit.TrackInfo + }{result1} +} + func (fake *FakeLocalParticipant) GetPlayoutDelayConfig() *livekit.PlayoutDelay { fake.getPlayoutDelayConfigMutex.Lock() ret, specificReturn := fake.getPlayoutDelayConfigReturnsOnCall[len(fake.getPlayoutDelayConfigArgsForCall)] @@ -5829,6 +5901,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.getLoggerMutex.RUnlock() fake.getPacerMutex.RLock() defer fake.getPacerMutex.RUnlock() + fake.getPendingTrackMutex.RLock() + defer fake.getPendingTrackMutex.RUnlock() fake.getPlayoutDelayConfigMutex.RLock() defer fake.getPlayoutDelayConfigMutex.RUnlock() fake.getPublishedTrackMutex.RLock()