diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 4491c6047..bf7420d28 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -4344,6 +4344,55 @@ func (p *ParticipantImpl) SupportsMoving() error { return nil } +// LeaveSession takes the participant out of its session in its current room without +// closing the participant, for it to join another session (see JoinSession): its +// published tracks are reported unpublished, the session end is reported, and the +// telemetry guard and deferred resolvers are reset for the next session. `leave` runs in +// between, while the left session's guard is still in place, for the caller to report the +// leave. +func (p *ParticipantImpl) LeaveSession(leave func()) { + for _, track := range p.GetPublishedTracks() { + p.GetTelemetryListener().OnTrackUnpublished( + p.ID(), + p.Identity(), + track.ToProto(), + track.(types.LocalMediaTrack).Published(), + true, + ) + } + + p.params.Reporter.ReportEndTime(time.Now()) + + if leave != nil { + leave() + } + + p.lock.Lock() + p.telemetryGuard = &telemetry.ReferenceGuard{} + p.lock.Unlock() + + p.params.LoggerResolver.Reset() + p.params.ReporterResolver.Reset() +} + +// JoinSession takes the participant, after LeaveSession, into the session its current +// room now serves: `join` runs first, for the caller to report the join, then the +// published tracks are reported published again, in that order as on a fresh join. +func (p *ParticipantImpl) JoinSession(join func()) { + if join != nil { + join() + } + + for _, track := range p.GetPublishedTracks() { + p.GetTelemetryListener().OnTrackPublished( + p.ID(), + p.Identity(), + track.ToProto(), + true, + ) + } +} + func (p *ParticipantImpl) MoveToRoom(params types.MoveToRoomParams) { for _, track := range p.GetPublishedTracks() { for _, sub := range track.GetAllSubscribers() { @@ -4353,38 +4402,23 @@ func (p *ParticipantImpl) MoveToRoom(params types.MoveToRoomParams) { // clear the subscriber node max quality/audio codecs as the remote quality notify // from source room would not reach the moving out participant. track.(types.LocalMediaTrack).ClearSubscriberNodes() - - trackInfo := track.ToProto() - p.GetTelemetryListener().OnTrackUnpublished( - p.ID(), - p.Identity(), - trackInfo, - track.(types.LocalMediaTrack).Published(), - true, - ) } - p.params.Reporter.ReportEndTime(time.Now()) - p.SubscriptionManager.ClearAllSubscriptions() + p.LeaveSession(func() { + p.SubscriptionManager.ClearAllSubscriptions() - // fire onClose callback for original room - p.lock.Lock() - onClose := p.onClose - p.onClose = make(map[string]func(types.LocalParticipant)) - p.lock.Unlock() - for _, cb := range onClose { - cb(p) - } + // fire onClose callback for original room + p.lock.Lock() + onClose := p.onClose + p.onClose = make(map[string]func(types.LocalParticipant)) + p.lock.Unlock() + for _, cb := range onClose { + cb(p) + } + }) p.params.Logger.Infow("move participant to new room", "newRoomName", params.RoomName, "newID", params.ParticipantID) - p.lock.Lock() - p.telemetryGuard = &telemetry.ReferenceGuard{} - p.lock.Unlock() - - p.params.LoggerResolver.Reset() - p.params.ReporterResolver.Reset() - p.setListener(params.Listener) p.setTelemetryListener(params.TelemetryListener) p.participantHelper.Store(params.Helper) diff --git a/pkg/rtc/participant_internal_test.go b/pkg/rtc/participant_internal_test.go index 7492beaf7..c2167eb53 100644 --- a/pkg/rtc/participant_internal_test.go +++ b/pkg/rtc/participant_internal_test.go @@ -32,6 +32,7 @@ import ( "github.com/livekit/protocol/codecs/mime" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" + "github.com/livekit/protocol/logger/zaputil" "github.com/livekit/protocol/observability/roomobs" lksdp "github.com/livekit/protocol/sdp" "github.com/livekit/protocol/signalling" @@ -1052,3 +1053,49 @@ func TestResumedParticipantWaitsForReconnectResponse(t *testing.T) { require.Equal(t, 1, sink.WriteMessageCallCount()) }) } + +func TestLeaveJoinSession(t *testing.T) { + p := newParticipantForTest("leave-join-session") + p.params.LoggerResolver = zaputil.NoOpDeferrer{} + _, p.params.ReporterResolver = roomobs.DeferredParticipantReporter(roomobs.NewNoopProjectReporter()) + tl := p.GetTelemetryListener().(*typesfakes.FakeParticipantTelemetryListener) + + track := &typesfakes.FakeLocalMediaTrack{} + track.IDReturns("TR_test") + track.PublishedReturns(true) + p.UpTrackManager.AddPublishedTrack(track) + + // leaving reports the track unpublished, then runs the leave under the left session's + // guard, and the next session gets a fresh one + prevGuard := p.TelemetryGuard() + require.NotNil(t, prevGuard) + left := false + p.LeaveSession(func() { + left = true + require.Same(t, prevGuard, p.TelemetryGuard()) + require.Equal(t, 1, tl.OnTrackUnpublishedCallCount()) + }) + require.True(t, left) + require.NotSame(t, prevGuard, p.TelemetryGuard()) + require.Equal(t, 1, tl.OnTrackUnpublishedCallCount()) + _, _, _, wasPublished, shouldSend := tl.OnTrackUnpublishedArgsForCall(0) + require.True(t, wasPublished) + require.True(t, shouldSend) + + // joining runs the join first, then reports the track published again + joined := false + p.JoinSession(func() { + joined = true + require.Zero(t, tl.OnTrackPublishedCallCount()) + }) + require.True(t, joined) + require.Equal(t, 1, tl.OnTrackPublishedCallCount()) + _, _, _, shouldSend = tl.OnTrackPublishedArgsForCall(0) + require.True(t, shouldSend) + + // nil callbacks are fine + p.LeaveSession(nil) + p.JoinSession(nil) + require.Equal(t, 2, tl.OnTrackUnpublishedCallCount()) + require.Equal(t, 2, tl.OnTrackPublishedCallCount()) +} diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index a8e3331df..1d7734f1d 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -608,6 +608,8 @@ type LocalParticipant interface { ) IsReconnect() bool IsMigration() bool + LeaveSession(leave func()) + JoinSession(join func()) MoveToRoom(params MoveToRoomParams) UpdateMediaRTT(rtt uint32) diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index a417a5b31..6c11e6bf6 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -1011,6 +1011,11 @@ type FakeLocalParticipant struct { issueFullReconnectArgsForCall []struct { arg1 types.ParticipantCloseReason } + JoinSessionStub func(func()) + joinSessionMutex sync.RWMutex + joinSessionArgsForCall []struct { + arg1 func() + } KindStub func() livekit.ParticipantInfo_Kind kindMutex sync.RWMutex kindArgsForCall []struct { @@ -1031,6 +1036,11 @@ type FakeLocalParticipant struct { kindDetailsReturnsOnCall map[int]struct { result1 []livekit.ParticipantInfo_KindDetail } + LeaveSessionStub func(func()) + leaveSessionMutex sync.RWMutex + leaveSessionArgsForCall []struct { + arg1 func() + } MaybeStartMigrationStub func(bool, func()) bool maybeStartMigrationMutex sync.RWMutex maybeStartMigrationArgsForCall []struct { @@ -6919,6 +6929,38 @@ func (fake *FakeLocalParticipant) IssueFullReconnectArgsForCall(i int) types.Par return argsForCall.arg1 } +func (fake *FakeLocalParticipant) JoinSession(arg1 func()) { + fake.joinSessionMutex.Lock() + fake.joinSessionArgsForCall = append(fake.joinSessionArgsForCall, struct { + arg1 func() + }{arg1}) + stub := fake.JoinSessionStub + fake.recordInvocation("JoinSession", []interface{}{arg1}) + fake.joinSessionMutex.Unlock() + if stub != nil { + fake.JoinSessionStub(arg1) + } +} + +func (fake *FakeLocalParticipant) JoinSessionCallCount() int { + fake.joinSessionMutex.RLock() + defer fake.joinSessionMutex.RUnlock() + return len(fake.joinSessionArgsForCall) +} + +func (fake *FakeLocalParticipant) JoinSessionCalls(stub func(func())) { + fake.joinSessionMutex.Lock() + defer fake.joinSessionMutex.Unlock() + fake.JoinSessionStub = stub +} + +func (fake *FakeLocalParticipant) JoinSessionArgsForCall(i int) func() { + fake.joinSessionMutex.RLock() + defer fake.joinSessionMutex.RUnlock() + argsForCall := fake.joinSessionArgsForCall[i] + return argsForCall.arg1 +} + func (fake *FakeLocalParticipant) Kind() livekit.ParticipantInfo_Kind { fake.kindMutex.Lock() ret, specificReturn := fake.kindReturnsOnCall[len(fake.kindArgsForCall)] @@ -7025,6 +7067,38 @@ func (fake *FakeLocalParticipant) KindDetailsReturnsOnCall(i int, result1 []live }{result1} } +func (fake *FakeLocalParticipant) LeaveSession(arg1 func()) { + fake.leaveSessionMutex.Lock() + fake.leaveSessionArgsForCall = append(fake.leaveSessionArgsForCall, struct { + arg1 func() + }{arg1}) + stub := fake.LeaveSessionStub + fake.recordInvocation("LeaveSession", []interface{}{arg1}) + fake.leaveSessionMutex.Unlock() + if stub != nil { + fake.LeaveSessionStub(arg1) + } +} + +func (fake *FakeLocalParticipant) LeaveSessionCallCount() int { + fake.leaveSessionMutex.RLock() + defer fake.leaveSessionMutex.RUnlock() + return len(fake.leaveSessionArgsForCall) +} + +func (fake *FakeLocalParticipant) LeaveSessionCalls(stub func(func())) { + fake.leaveSessionMutex.Lock() + defer fake.leaveSessionMutex.Unlock() + fake.LeaveSessionStub = stub +} + +func (fake *FakeLocalParticipant) LeaveSessionArgsForCall(i int) func() { + fake.leaveSessionMutex.RLock() + defer fake.leaveSessionMutex.RUnlock() + argsForCall := fake.leaveSessionArgsForCall[i] + return argsForCall.arg1 +} + func (fake *FakeLocalParticipant) MaybeStartMigration(arg1 bool, arg2 func()) bool { fake.maybeStartMigrationMutex.Lock() ret, specificReturn := fake.maybeStartMigrationReturnsOnCall[len(fake.maybeStartMigrationArgsForCall)]