diff --git a/pkg/rtc/datatrack.go b/pkg/rtc/datatrack.go index e68c7efd9..540e69bf2 100644 --- a/pkg/rtc/datatrack.go +++ b/pkg/rtc/datatrack.go @@ -88,6 +88,12 @@ func (t *DataTrack) OnClose(f func()) { t.onClose = f } +func (t *DataTrack) IsSubscriber(subId string) bool { + t.lock.RLock() + defer t.lock.RUnlock() + return t.subscribers[subId] != nil +} + func (t *DataTrack) AddSubscriber(participant types.Participant) error { t.lock.Lock() defer t.lock.Unlock() diff --git a/pkg/rtc/mediatrack.go b/pkg/rtc/mediatrack.go index 29b6a445b..fee799166 100644 --- a/pkg/rtc/mediatrack.go +++ b/pkg/rtc/mediatrack.go @@ -107,6 +107,12 @@ func (t *MediaTrack) OnClose(f func()) { t.onClose = f } +func (t *MediaTrack) IsSubscriber(subId string) bool { + t.lock.RLock() + defer t.lock.RUnlock() + return t.subscribedTracks[subId] != nil +} + // AddSubscriber subscribes sub to current mediaTrack func (t *MediaTrack) AddSubscriber(sub types.Participant) error { t.lock.Lock() diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 204e6c425..4e1b4464b 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -88,6 +88,7 @@ type PublishedTrack interface { SetMuted(muted bool) AddSubscriber(participant Participant) error RemoveSubscriber(participantId string) + IsSubscriber(subId string) bool RemoveAllSubscribers() // callbacks diff --git a/pkg/rtc/types/typesfakes/fake_published_track.go b/pkg/rtc/types/typesfakes/fake_published_track.go index fb1b55fcf..b0aaef18b 100644 --- a/pkg/rtc/types/typesfakes/fake_published_track.go +++ b/pkg/rtc/types/typesfakes/fake_published_track.go @@ -40,6 +40,17 @@ type FakePublishedTrack struct { isMutedReturnsOnCall map[int]struct { result1 bool } + IsSubscriberStub func(string) bool + isSubscriberMutex sync.RWMutex + isSubscriberArgsForCall []struct { + arg1 string + } + isSubscriberReturns struct { + result1 bool + } + isSubscriberReturnsOnCall map[int]struct { + result1 bool + } KindStub func() livekit.TrackType kindMutex sync.RWMutex kindArgsForCall []struct { @@ -254,6 +265,67 @@ func (fake *FakePublishedTrack) IsMutedReturnsOnCall(i int, result1 bool) { }{result1} } +func (fake *FakePublishedTrack) IsSubscriber(arg1 string) bool { + fake.isSubscriberMutex.Lock() + ret, specificReturn := fake.isSubscriberReturnsOnCall[len(fake.isSubscriberArgsForCall)] + fake.isSubscriberArgsForCall = append(fake.isSubscriberArgsForCall, struct { + arg1 string + }{arg1}) + stub := fake.IsSubscriberStub + fakeReturns := fake.isSubscriberReturns + fake.recordInvocation("IsSubscriber", []interface{}{arg1}) + fake.isSubscriberMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakePublishedTrack) IsSubscriberCallCount() int { + fake.isSubscriberMutex.RLock() + defer fake.isSubscriberMutex.RUnlock() + return len(fake.isSubscriberArgsForCall) +} + +func (fake *FakePublishedTrack) IsSubscriberCalls(stub func(string) bool) { + fake.isSubscriberMutex.Lock() + defer fake.isSubscriberMutex.Unlock() + fake.IsSubscriberStub = stub +} + +func (fake *FakePublishedTrack) IsSubscriberArgsForCall(i int) string { + fake.isSubscriberMutex.RLock() + defer fake.isSubscriberMutex.RUnlock() + argsForCall := fake.isSubscriberArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakePublishedTrack) IsSubscriberReturns(result1 bool) { + fake.isSubscriberMutex.Lock() + defer fake.isSubscriberMutex.Unlock() + fake.IsSubscriberStub = nil + fake.isSubscriberReturns = struct { + result1 bool + }{result1} +} + +func (fake *FakePublishedTrack) IsSubscriberReturnsOnCall(i int, result1 bool) { + fake.isSubscriberMutex.Lock() + defer fake.isSubscriberMutex.Unlock() + fake.IsSubscriberStub = nil + if fake.isSubscriberReturnsOnCall == nil { + fake.isSubscriberReturnsOnCall = make(map[int]struct { + result1 bool + }) + } + fake.isSubscriberReturnsOnCall[i] = struct { + result1 bool + }{result1} +} + func (fake *FakePublishedTrack) Kind() livekit.TrackType { fake.kindMutex.Lock() ret, specificReturn := fake.kindReturnsOnCall[len(fake.kindArgsForCall)] @@ -513,6 +585,8 @@ func (fake *FakePublishedTrack) Invocations() map[string][][]interface{} { defer fake.iDMutex.RUnlock() fake.isMutedMutex.RLock() defer fake.isMutedMutex.RUnlock() + fake.isSubscriberMutex.RLock() + defer fake.isSubscriberMutex.RUnlock() fake.kindMutex.RLock() defer fake.kindMutex.RUnlock() fake.nameMutex.RLock() diff --git a/pkg/service/roommanager.go b/pkg/service/roommanager.go index 076a612a5..b471602c2 100644 --- a/pkg/service/roommanager.go +++ b/pkg/service/roommanager.go @@ -111,6 +111,12 @@ func (r *RoomManager) CreateRoom(req *livekit.CreateRoomRequest) (*livekit.Room, return rm, nil } +func (r *RoomManager) GetRoom(roomName string) *rtc.Room { + r.lock.RLock() + defer r.lock.RUnlock() + return r.rooms[roomName] +} + // DeleteRoom completely deletes all room information, including active sessions, room store, and routing info func (r *RoomManager) DeleteRoom(roomName string) error { logger.Infow("deleting room state", "room", roomName) diff --git a/pkg/service/server.go b/pkg/service/server.go index 837b171e3..30471cb4c 100644 --- a/pkg/service/server.go +++ b/pkg/service/server.go @@ -196,6 +196,10 @@ func (s *LivekitServer) Stop() { } } +func (s *LivekitServer) RoomManager() *RoomManager { + return s.roomManager +} + func (s *LivekitServer) healthCheck(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) } diff --git a/test/singlenode_test.go b/test/singlenode_test.go index f1d6b66da..9d3082dea 100644 --- a/test/singlenode_test.go +++ b/test/singlenode_test.go @@ -5,6 +5,7 @@ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestClientCouldConnect(t *testing.T) { @@ -96,4 +97,22 @@ func TestSinglePublisher(t *testing.T) { for _, tr := range tracks { assert.True(t, strings.HasPrefix(tr.ID(), "TR_"), "track should begin with TR") } + + // when c3 disconnects.. ensure subscriber is cleaned up correctly + c3.Stop() + + success = withTimeout(t, "c3 is cleaned up as a subscriber", func() bool { + room := s.RoomManager().GetRoom(testRoom) + require.NotNil(t, room) + + p := room.GetParticipant("c3") + require.NotNil(t, p) + + for _, t := range p.GetPublishedTracks() { + if t.IsSubscriber(p.ID()) { + return false + } + } + return true + }) }