From 913ef3a6467d98e36d6ee3854fea115c70fa18d7 Mon Sep 17 00:00:00 2001 From: cnderrauber Date: Tue, 1 Mar 2022 15:48:20 +0800 Subject: [PATCH] Datatrack for data channel (#476) * data track --- pkg/rtc/datatrack.go | 189 +++++++++++ pkg/rtc/participant.go | 89 +++-- pkg/rtc/room.go | 8 +- pkg/rtc/room_test.go | 12 +- pkg/rtc/types/interfaces.go | 13 +- pkg/rtc/types/typesfakes/fake_data_track.go | 312 ++++++++++++++++++ .../typesfakes/fake_local_participant.go | 117 +++++-- pkg/rtc/types/typesfakes/fake_participant.go | 65 ++++ pkg/sfu/receiver.go | 15 +- test/multinode_roomservice_test.go | 4 +- test/singlenode_test.go | 2 +- 11 files changed, 738 insertions(+), 88 deletions(-) create mode 100644 pkg/rtc/datatrack.go create mode 100644 pkg/rtc/types/typesfakes/fake_data_track.go diff --git a/pkg/rtc/datatrack.go b/pkg/rtc/datatrack.go new file mode 100644 index 000000000..818ae2e1b --- /dev/null +++ b/pkg/rtc/datatrack.go @@ -0,0 +1,189 @@ +package rtc + +import ( + "errors" + "sync" + + "github.com/livekit/livekit-server/pkg/sfu" + "github.com/livekit/protocol/livekit" + "github.com/livekit/protocol/logger" + "github.com/pion/webrtc/v3" + "google.golang.org/protobuf/proto" +) + +type DataTrackSender interface { + sfu.TrackSender + Write(label string, data []byte) +} + +type DataTrack struct { + trackID livekit.TrackID + participantID livekit.ParticipantID + logger logger.Logger + lock sync.RWMutex + downTracks []DataTrackSender + onDataPacket func(*livekit.DataPacket) + onClose []func() +} + +func NewDataTrack(trackID livekit.TrackID, participantID livekit.ParticipantID, logger logger.Logger) *DataTrack { + t := &DataTrack{ + trackID: trackID, + participantID: participantID, + logger: logger, + } + return t +} + +func (t *DataTrack) onData(label string, data []byte) { + t.lock.RLock() + f := t.onDataPacket + dts := t.downTracks + t.lock.RUnlock() + + for _, dt := range dts { + dt.Write(label, data) + } + + if f != nil { + dp, err := DataPacketFromBytes(label, data) + if err != nil { + t.logger.Warnw("invalid data", err, "label", label) + return + } + // only forward on user payloads + switch payload := dp.Value.(type) { + case *livekit.DataPacket_User: + payload.User.ParticipantSid = string(t.participantID) + f(dp) + default: + t.logger.Warnw("received unsupported data packet", nil, "payload", payload) + } + } +} + +func (t *DataTrack) OnDataPacket(f func(*livekit.DataPacket)) { + t.lock.Lock() + t.onDataPacket = f + t.lock.Unlock() +} + +func (t *DataTrack) TrackID() livekit.TrackID { + return t.trackID +} + +func (t *DataTrack) Write(label string, data []byte) { + t.onData(label, data) +} + +func (t *DataTrack) AddDownTrack(dt sfu.TrackSender) error { + dataDt, ok := dt.(DataTrackSender) + if !ok { + return errors.New("invalid DownTrack type, expect DataTrackSender") + } + t.lock.Lock() + defer t.lock.Unlock() + t.downTracks = append(t.downTracks, dataDt) + return nil +} + +func (t *DataTrack) DeleteDownTrack(peerID livekit.ParticipantID) { + t.lock.Lock() + defer t.lock.Unlock() + for k, v := range t.downTracks { + if v.PeerID() == peerID { + t.downTracks[k] = t.downTracks[len(t.downTracks)-1] + t.downTracks = t.downTracks[:len(t.downTracks)-1] + break + } + } +} + +func (t *DataTrack) AddOnClose(f func()) { + if f == nil { + return + } + t.lock.Lock() + t.onClose = append(t.onClose, f) + t.lock.Unlock() +} + +func (t *DataTrack) Close() { + t.lock.Lock() + fs := t.onClose + t.lock.Unlock() + + for _, f := range fs { + f() + } +} + +func (t *DataTrack) Receiver() sfu.TrackReceiver { + return t +} + +func (t *DataTrack) ToProto() *livekit.TrackInfo { + return &livekit.TrackInfo{ + Sid: string(t.trackID), + Type: livekit.TrackType_DATA, + } +} + +func DataPacketFromBytes(label string, data []byte) (*livekit.DataPacket, error) { + dp := livekit.DataPacket{} + if err := proto.Unmarshal(data, &dp); err != nil { + return nil, err + } + + switch label { + case reliableDataChannel: + dp.Kind = livekit.DataPacket_RELIABLE + case lossyDataChannel: + dp.Kind = livekit.DataPacket_LOSSY + default: + return nil, errors.New("unsupported datachannel added") + } + + return &dp, nil +} + +func (t *DataTrack) Kind() livekit.TrackType { + return livekit.TrackType_DATA +} + +//--------------------------------------------- +// no op methods for sfu.TrackReceiver +func (t *DataTrack) StreamID() string { + return "" +} + +func (t *DataTrack) Codec() webrtc.RTPCodecCapability { + return webrtc.RTPCodecCapability{} +} + +func (t *DataTrack) ReadRTP(buf []byte, layer uint8, sn uint16) (int, error) { + return 0, nil +} + +func (t *DataTrack) GetSenderReportTime(layer int32) (rtpTS uint32, ntpTS uint64) { + return +} + +func (t *DataTrack) GetBitrateTemporalCumulative() sfu.Bitrates { + return sfu.Bitrates{} +} + +func (t *DataTrack) SendPLI(layer int32) { +} + +func (t *DataTrack) SetUpTrackPaused(paused bool) { + +} + +func (t *DataTrack) SetMaxExpectedSpatialLayer(layer int32) { + +} + +func (t *DataTrack) DebugInfo() map[string]interface{} { + return map[string]interface{}{} +} diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index fa1eef8d4..71d32ddd8 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -109,23 +109,25 @@ type ParticipantImpl struct { updateLock sync.Mutex version uint32 + dataTrack *DataTrack + // callbacks & handlers onTrackPublished func(types.LocalParticipant, types.MediaTrack) onTrackUpdated func(types.LocalParticipant, types.MediaTrack) onStateChange func(p types.LocalParticipant, oldState livekit.ParticipantInfo_State) onMetadataUpdate func(types.LocalParticipant) - onDataPacket func(types.LocalParticipant, *livekit.DataPacket) migrateState atomic.Value // types.MigrateState pendingOffer *webrtc.SessionDescription pendingDataChannels []*livekit.DataChannelInfo onClose func(types.LocalParticipant, map[livekit.TrackID]livekit.ParticipantID) onClaimsChanged func(participant types.LocalParticipant) + + onDataTrackPublished func(types.LocalParticipant, types.DataTrack) } func NewParticipant(params ParticipantParams, perms *livekit.ParticipantPermission) (*ParticipantImpl, error) { // TODO: check to ensure params are valid, id and identity can't be empty - p := &ParticipantImpl{ params: params, rtcpCh: make(chan []rtcp.Packet, 50), @@ -312,6 +314,10 @@ func (p *ParticipantImpl) ToProto() *livekit.ParticipantInfo { info.Metadata = p.params.Grants.Metadata } + if p.dataTrack != nil { + info.Tracks = append(info.Tracks, p.dataTrack.ToProto()) + } + return info } @@ -345,10 +351,6 @@ func (p *ParticipantImpl) OnMetadataUpdate(callback func(types.LocalParticipant) p.onMetadataUpdate = callback } -func (p *ParticipantImpl) OnDataPacket(callback func(types.LocalParticipant, *livekit.DataPacket)) { - p.onDataPacket = callback -} - func (p *ParticipantImpl) OnClose(callback func(types.LocalParticipant, map[livekit.TrackID]livekit.ParticipantID)) { p.onClose = callback } @@ -496,6 +498,9 @@ func (p *ParticipantImpl) Close(sendLeave bool) error { }) } + if p.dataTrack != nil { + p.dataTrack.Close() + } p.UpTrackManager.Close() p.pendingTracksLock.Lock() @@ -1071,48 +1076,37 @@ func (p *ParticipantImpl) onMediaTrack(track *webrtc.TrackRemote, rtpReceiver *w } } +func (p *ParticipantImpl) OnDataTrackPublished(f func(types.LocalParticipant, types.DataTrack)) { + p.onDataTrackPublished = f +} + func (p *ParticipantImpl) onDataChannel(dc *webrtc.DataChannel) { if p.State() == livekit.ParticipantInfo_DISCONNECTED { return } - switch dc.Label() { + if p.dataTrack == nil { + p.dataTrack = NewDataTrack(livekit.TrackID("DT_"+p.params.SID), p.params.SID, p.params.Logger) + if p.onDataTrackPublished != nil { + p.onDataTrackPublished(p, p.dataTrack) + } + } + label := dc.Label() + switch label { case reliableDataChannel: p.reliableDC = dc dc.OnMessage(func(msg webrtc.DataChannelMessage) { - p.handleDataMessage(livekit.DataPacket_RELIABLE, msg.Data) + p.dataTrack.Write(label, msg.Data) }) case lossyDataChannel: p.lossyDC = dc dc.OnMessage(func(msg webrtc.DataChannelMessage) { - p.handleDataMessage(livekit.DataPacket_LOSSY, msg.Data) + p.dataTrack.Write(label, msg.Data) }) default: p.params.Logger.Warnw("unsupported datachannel added", nil, "label", dc.Label()) } } -func (p *ParticipantImpl) handleDataMessage(kind livekit.DataPacket_Kind, data []byte) { - dp := livekit.DataPacket{} - if err := proto.Unmarshal(data, &dp); err != nil { - p.params.Logger.Warnw("could not parse data packet", err) - return - } - - // trust the channel that it came in as the source of truth - dp.Kind = kind - - // only forward on user payloads - switch payload := dp.Value.(type) { - case *livekit.DataPacket_User: - if p.onDataPacket != nil { - payload.User.ParticipantSid = string(p.params.SID) - p.onDataPacket(p, &dp) - } - default: - p.params.Logger.Warnw("received unsupported data packet", nil, "payload", payload) - } -} - func (p *ParticipantImpl) handlePrimaryStateChange(state webrtc.PeerConnectionState) { if state == webrtc.PeerConnectionStateConnected { prometheus.ServiceOperationCounter.WithLabelValues("ice_connection", "success", "").Add(1) @@ -1613,44 +1607,41 @@ func (p *ParticipantImpl) DebugInfo() map[string]interface{} { return info } +func (p *ParticipantImpl) GetDataTrack() types.DataTrack { + return p.dataTrack +} + func (p *ParticipantImpl) handlePendingDataChannels() { p.lock.Lock() defer p.lock.Unlock() ordered := true negotiated := true for _, ci := range p.pendingDataChannels { + var ( + dc *webrtc.DataChannel + err error + ) if ci.Label == lossyDataChannel && p.lossyDC == nil { retransmits := uint16(0) id := uint16(ci.GetId()) - dc, err := p.publisher.pc.CreateDataChannel(lossyDataChannel, &webrtc.DataChannelInit{ + dc, err = p.publisher.pc.CreateDataChannel(lossyDataChannel, &webrtc.DataChannelInit{ Ordered: &ordered, MaxRetransmits: &retransmits, Negotiated: &negotiated, ID: &id, }) - if err != nil { - p.params.Logger.Errorw("create migrated data channel failed", err, "label", lossyDataChannel) - } else { - p.lossyDC = dc - dc.OnMessage(func(msg webrtc.DataChannelMessage) { - p.handleDataMessage(livekit.DataPacket_LOSSY, msg.Data) - }) - } } else if ci.Label == reliableDataChannel && p.reliableDC == nil { id := uint16(ci.GetId()) - dc, err := p.publisher.pc.CreateDataChannel(reliableDataChannel, &webrtc.DataChannelInit{ + dc, err = p.publisher.pc.CreateDataChannel(reliableDataChannel, &webrtc.DataChannelInit{ Ordered: &ordered, Negotiated: &negotiated, ID: &id, }) - if err != nil { - p.params.Logger.Errorw("create migrated data channel failed", err, "label", reliableDataChannel) - } else { - p.reliableDC = dc - dc.OnMessage(func(msg webrtc.DataChannelMessage) { - p.handleDataMessage(livekit.DataPacket_RELIABLE, msg.Data) - }) - } + } + if err != nil { + p.params.Logger.Errorw("create migrated data channel failed", err, "label", ci.Label) + } else if dc != nil { + p.onDataChannel(dc) } } p.pendingDataChannels = nil diff --git a/pkg/rtc/room.go b/pkg/rtc/room.go index 0d6dfaff6..10d06f793 100644 --- a/pkg/rtc/room.go +++ b/pkg/rtc/room.go @@ -221,7 +221,11 @@ func (r *Room) Join(participant types.LocalParticipant, opts *ParticipantOptions }) participant.OnTrackUpdated(r.onTrackUpdated) participant.OnMetadataUpdate(r.onParticipantMetadataUpdate) - participant.OnDataPacket(r.onDataPacket) + participant.OnDataTrackPublished(func(lp types.LocalParticipant, dt types.DataTrack) { + dt.OnDataPacket(func(dp *livekit.DataPacket) { + r.onDataPacket(lp, dp) + }) + }) r.Logger.Infow("new participant joined", "pID", participant.ID(), "participant", participant.Identity(), @@ -328,7 +332,7 @@ func (r *Room) RemoveParticipant(identity livekit.ParticipantIdentity) { p.OnTrackPublished(nil) p.OnStateChange(nil) p.OnMetadataUpdate(nil) - p.OnDataPacket(nil) + p.OnDataTrackPublished(nil) // close participant as well _ = p.Close(true) diff --git a/pkg/rtc/room_test.go b/pkg/rtc/room_test.go index e683b7c7a..d62e5dbc0 100644 --- a/pkg/rtc/room_test.go +++ b/pkg/rtc/room_test.go @@ -420,7 +420,9 @@ func TestDataChannel(t *testing.T) { }, }, } - p.OnDataPacketArgsForCall(0)(p, &packet) + dataTrack := &typesfakes.FakeDataTrack{} + p.OnDataTrackPublishedArgsForCall(0)(p, dataTrack) + dataTrack.OnDataPacketArgsForCall(0)(&packet) // ensure everyone has received the packet for _, op := range participants { @@ -451,7 +453,9 @@ func TestDataChannel(t *testing.T) { }, }, } - p.OnDataPacketArgsForCall(0)(p, &packet) + dataTrack := &typesfakes.FakeDataTrack{} + p.OnDataTrackPublishedArgsForCall(0)(p, dataTrack) + dataTrack.OnDataPacketArgsForCall(0)(&packet) // only p1 should receive the data for _, op := range participants { @@ -479,7 +483,9 @@ func TestDataChannel(t *testing.T) { }, }, } - p.OnDataPacketArgsForCall(0)(p, &packet) + dataTrack := &typesfakes.FakeDataTrack{} + p.OnDataTrackPublishedArgsForCall(0)(p, dataTrack) + dataTrack.OnDataPacketArgsForCall(0)(&packet) // no one should've been sent packet for _, op := range participants { diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 87964fd69..564570c9a 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -59,6 +59,7 @@ type Participant interface { GetPublishedTrack(sid livekit.TrackID) MediaTrack GetPublishedTracks() []MediaTrack + GetDataTrack() DataTrack AddSubscriber(op LocalParticipant, params AddSubscriberParams) (int, error) RemoveSubscriber(op LocalParticipant, trackID livekit.TrackID, resume bool) @@ -144,9 +145,9 @@ type LocalParticipant interface { // OnTrackUpdated - one of its publishedTracks changed in status OnTrackUpdated(callback func(LocalParticipant, MediaTrack)) OnMetadataUpdate(callback func(LocalParticipant)) - OnDataPacket(callback func(LocalParticipant, *livekit.DataPacket)) OnClose(_callback func(LocalParticipant, map[livekit.TrackID]livekit.ParticipantID)) OnClaimsChanged(_callback func(LocalParticipant)) + OnDataTrackPublished(callback func(LocalParticipant, DataTrack)) // session migration SetMigrateState(s MigrateState) @@ -237,3 +238,13 @@ type SubscribedTrack interface { // selects appropriate video layer according to subscriber preferences UpdateVideoLayer() } + +// DataTrack is the interface representing a data track published to the room +//counterfeiter:generate . DataTrack +type DataTrack interface { + TrackID() livekit.TrackID + Receiver() sfu.TrackReceiver + AddOnClose(func()) + OnDataPacket(callback func(*livekit.DataPacket)) + Kind() livekit.TrackType +} diff --git a/pkg/rtc/types/typesfakes/fake_data_track.go b/pkg/rtc/types/typesfakes/fake_data_track.go new file mode 100644 index 000000000..afedc7506 --- /dev/null +++ b/pkg/rtc/types/typesfakes/fake_data_track.go @@ -0,0 +1,312 @@ +// Code generated by counterfeiter. DO NOT EDIT. +package typesfakes + +import ( + "sync" + + "github.com/livekit/livekit-server/pkg/rtc/types" + "github.com/livekit/livekit-server/pkg/sfu" + "github.com/livekit/protocol/livekit" +) + +type FakeDataTrack struct { + AddOnCloseStub func(func()) + addOnCloseMutex sync.RWMutex + addOnCloseArgsForCall []struct { + arg1 func() + } + KindStub func() livekit.TrackType + kindMutex sync.RWMutex + kindArgsForCall []struct { + } + kindReturns struct { + result1 livekit.TrackType + } + kindReturnsOnCall map[int]struct { + result1 livekit.TrackType + } + OnDataPacketStub func(func(*livekit.DataPacket)) + onDataPacketMutex sync.RWMutex + onDataPacketArgsForCall []struct { + arg1 func(*livekit.DataPacket) + } + ReceiverStub func() sfu.TrackReceiver + receiverMutex sync.RWMutex + receiverArgsForCall []struct { + } + receiverReturns struct { + result1 sfu.TrackReceiver + } + receiverReturnsOnCall map[int]struct { + result1 sfu.TrackReceiver + } + TrackIDStub func() livekit.TrackID + trackIDMutex sync.RWMutex + trackIDArgsForCall []struct { + } + trackIDReturns struct { + result1 livekit.TrackID + } + trackIDReturnsOnCall map[int]struct { + result1 livekit.TrackID + } + invocations map[string][][]interface{} + invocationsMutex sync.RWMutex +} + +func (fake *FakeDataTrack) AddOnClose(arg1 func()) { + fake.addOnCloseMutex.Lock() + fake.addOnCloseArgsForCall = append(fake.addOnCloseArgsForCall, struct { + arg1 func() + }{arg1}) + stub := fake.AddOnCloseStub + fake.recordInvocation("AddOnClose", []interface{}{arg1}) + fake.addOnCloseMutex.Unlock() + if stub != nil { + fake.AddOnCloseStub(arg1) + } +} + +func (fake *FakeDataTrack) AddOnCloseCallCount() int { + fake.addOnCloseMutex.RLock() + defer fake.addOnCloseMutex.RUnlock() + return len(fake.addOnCloseArgsForCall) +} + +func (fake *FakeDataTrack) AddOnCloseCalls(stub func(func())) { + fake.addOnCloseMutex.Lock() + defer fake.addOnCloseMutex.Unlock() + fake.AddOnCloseStub = stub +} + +func (fake *FakeDataTrack) AddOnCloseArgsForCall(i int) func() { + fake.addOnCloseMutex.RLock() + defer fake.addOnCloseMutex.RUnlock() + argsForCall := fake.addOnCloseArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakeDataTrack) Kind() livekit.TrackType { + fake.kindMutex.Lock() + ret, specificReturn := fake.kindReturnsOnCall[len(fake.kindArgsForCall)] + fake.kindArgsForCall = append(fake.kindArgsForCall, struct { + }{}) + stub := fake.KindStub + fakeReturns := fake.kindReturns + fake.recordInvocation("Kind", []interface{}{}) + fake.kindMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeDataTrack) KindCallCount() int { + fake.kindMutex.RLock() + defer fake.kindMutex.RUnlock() + return len(fake.kindArgsForCall) +} + +func (fake *FakeDataTrack) KindCalls(stub func() livekit.TrackType) { + fake.kindMutex.Lock() + defer fake.kindMutex.Unlock() + fake.KindStub = stub +} + +func (fake *FakeDataTrack) KindReturns(result1 livekit.TrackType) { + fake.kindMutex.Lock() + defer fake.kindMutex.Unlock() + fake.KindStub = nil + fake.kindReturns = struct { + result1 livekit.TrackType + }{result1} +} + +func (fake *FakeDataTrack) KindReturnsOnCall(i int, result1 livekit.TrackType) { + fake.kindMutex.Lock() + defer fake.kindMutex.Unlock() + fake.KindStub = nil + if fake.kindReturnsOnCall == nil { + fake.kindReturnsOnCall = make(map[int]struct { + result1 livekit.TrackType + }) + } + fake.kindReturnsOnCall[i] = struct { + result1 livekit.TrackType + }{result1} +} + +func (fake *FakeDataTrack) OnDataPacket(arg1 func(*livekit.DataPacket)) { + fake.onDataPacketMutex.Lock() + fake.onDataPacketArgsForCall = append(fake.onDataPacketArgsForCall, struct { + arg1 func(*livekit.DataPacket) + }{arg1}) + stub := fake.OnDataPacketStub + fake.recordInvocation("OnDataPacket", []interface{}{arg1}) + fake.onDataPacketMutex.Unlock() + if stub != nil { + fake.OnDataPacketStub(arg1) + } +} + +func (fake *FakeDataTrack) OnDataPacketCallCount() int { + fake.onDataPacketMutex.RLock() + defer fake.onDataPacketMutex.RUnlock() + return len(fake.onDataPacketArgsForCall) +} + +func (fake *FakeDataTrack) OnDataPacketCalls(stub func(func(*livekit.DataPacket))) { + fake.onDataPacketMutex.Lock() + defer fake.onDataPacketMutex.Unlock() + fake.OnDataPacketStub = stub +} + +func (fake *FakeDataTrack) OnDataPacketArgsForCall(i int) func(*livekit.DataPacket) { + fake.onDataPacketMutex.RLock() + defer fake.onDataPacketMutex.RUnlock() + argsForCall := fake.onDataPacketArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakeDataTrack) Receiver() sfu.TrackReceiver { + fake.receiverMutex.Lock() + ret, specificReturn := fake.receiverReturnsOnCall[len(fake.receiverArgsForCall)] + fake.receiverArgsForCall = append(fake.receiverArgsForCall, struct { + }{}) + stub := fake.ReceiverStub + fakeReturns := fake.receiverReturns + fake.recordInvocation("Receiver", []interface{}{}) + fake.receiverMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeDataTrack) ReceiverCallCount() int { + fake.receiverMutex.RLock() + defer fake.receiverMutex.RUnlock() + return len(fake.receiverArgsForCall) +} + +func (fake *FakeDataTrack) ReceiverCalls(stub func() sfu.TrackReceiver) { + fake.receiverMutex.Lock() + defer fake.receiverMutex.Unlock() + fake.ReceiverStub = stub +} + +func (fake *FakeDataTrack) ReceiverReturns(result1 sfu.TrackReceiver) { + fake.receiverMutex.Lock() + defer fake.receiverMutex.Unlock() + fake.ReceiverStub = nil + fake.receiverReturns = struct { + result1 sfu.TrackReceiver + }{result1} +} + +func (fake *FakeDataTrack) ReceiverReturnsOnCall(i int, result1 sfu.TrackReceiver) { + fake.receiverMutex.Lock() + defer fake.receiverMutex.Unlock() + fake.ReceiverStub = nil + if fake.receiverReturnsOnCall == nil { + fake.receiverReturnsOnCall = make(map[int]struct { + result1 sfu.TrackReceiver + }) + } + fake.receiverReturnsOnCall[i] = struct { + result1 sfu.TrackReceiver + }{result1} +} + +func (fake *FakeDataTrack) TrackID() livekit.TrackID { + fake.trackIDMutex.Lock() + ret, specificReturn := fake.trackIDReturnsOnCall[len(fake.trackIDArgsForCall)] + fake.trackIDArgsForCall = append(fake.trackIDArgsForCall, struct { + }{}) + stub := fake.TrackIDStub + fakeReturns := fake.trackIDReturns + fake.recordInvocation("TrackID", []interface{}{}) + fake.trackIDMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeDataTrack) TrackIDCallCount() int { + fake.trackIDMutex.RLock() + defer fake.trackIDMutex.RUnlock() + return len(fake.trackIDArgsForCall) +} + +func (fake *FakeDataTrack) TrackIDCalls(stub func() livekit.TrackID) { + fake.trackIDMutex.Lock() + defer fake.trackIDMutex.Unlock() + fake.TrackIDStub = stub +} + +func (fake *FakeDataTrack) TrackIDReturns(result1 livekit.TrackID) { + fake.trackIDMutex.Lock() + defer fake.trackIDMutex.Unlock() + fake.TrackIDStub = nil + fake.trackIDReturns = struct { + result1 livekit.TrackID + }{result1} +} + +func (fake *FakeDataTrack) TrackIDReturnsOnCall(i int, result1 livekit.TrackID) { + fake.trackIDMutex.Lock() + defer fake.trackIDMutex.Unlock() + fake.TrackIDStub = nil + if fake.trackIDReturnsOnCall == nil { + fake.trackIDReturnsOnCall = make(map[int]struct { + result1 livekit.TrackID + }) + } + fake.trackIDReturnsOnCall[i] = struct { + result1 livekit.TrackID + }{result1} +} + +func (fake *FakeDataTrack) Invocations() map[string][][]interface{} { + fake.invocationsMutex.RLock() + defer fake.invocationsMutex.RUnlock() + fake.addOnCloseMutex.RLock() + defer fake.addOnCloseMutex.RUnlock() + fake.kindMutex.RLock() + defer fake.kindMutex.RUnlock() + fake.onDataPacketMutex.RLock() + defer fake.onDataPacketMutex.RUnlock() + fake.receiverMutex.RLock() + defer fake.receiverMutex.RUnlock() + fake.trackIDMutex.RLock() + defer fake.trackIDMutex.RUnlock() + copiedInvocations := map[string][][]interface{}{} + for key, value := range fake.invocations { + copiedInvocations[key] = value + } + return copiedInvocations +} + +func (fake *FakeDataTrack) recordInvocation(key string, args []interface{}) { + fake.invocationsMutex.Lock() + defer fake.invocationsMutex.Unlock() + if fake.invocations == nil { + fake.invocations = map[string][][]interface{}{} + } + if fake.invocations[key] == nil { + fake.invocations[key] = [][]interface{}{} + } + fake.invocations[key] = append(fake.invocations[key], args) +} + +var _ types.DataTrack = new(FakeDataTrack) diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index 1a71a2145..f68acba81 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -143,6 +143,16 @@ type FakeLocalParticipant struct { getConnectionQualityReturnsOnCall map[int]struct { result1 *livekit.ConnectionQualityInfo } + GetDataTrackStub func() types.DataTrack + getDataTrackMutex sync.RWMutex + getDataTrackArgsForCall []struct { + } + getDataTrackReturns struct { + result1 types.DataTrack + } + getDataTrackReturnsOnCall map[int]struct { + result1 types.DataTrack + } GetLoggerStub func() logger.Logger getLoggerMutex sync.RWMutex getLoggerArgsForCall []struct { @@ -302,10 +312,10 @@ type FakeLocalParticipant struct { onCloseArgsForCall []struct { arg1 func(types.LocalParticipant, map[livekit.TrackID]livekit.ParticipantID) } - OnDataPacketStub func(func(types.LocalParticipant, *livekit.DataPacket)) - onDataPacketMutex sync.RWMutex - onDataPacketArgsForCall []struct { - arg1 func(types.LocalParticipant, *livekit.DataPacket) + OnDataTrackPublishedStub func(func(types.LocalParticipant, types.DataTrack)) + onDataTrackPublishedMutex sync.RWMutex + onDataTrackPublishedArgsForCall []struct { + arg1 func(types.LocalParticipant, types.DataTrack) } OnMetadataUpdateStub func(func(types.LocalParticipant)) onMetadataUpdateMutex sync.RWMutex @@ -1286,6 +1296,59 @@ func (fake *FakeLocalParticipant) GetConnectionQualityReturnsOnCall(i int, resul }{result1} } +func (fake *FakeLocalParticipant) GetDataTrack() types.DataTrack { + fake.getDataTrackMutex.Lock() + ret, specificReturn := fake.getDataTrackReturnsOnCall[len(fake.getDataTrackArgsForCall)] + fake.getDataTrackArgsForCall = append(fake.getDataTrackArgsForCall, struct { + }{}) + stub := fake.GetDataTrackStub + fakeReturns := fake.getDataTrackReturns + fake.recordInvocation("GetDataTrack", []interface{}{}) + fake.getDataTrackMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeLocalParticipant) GetDataTrackCallCount() int { + fake.getDataTrackMutex.RLock() + defer fake.getDataTrackMutex.RUnlock() + return len(fake.getDataTrackArgsForCall) +} + +func (fake *FakeLocalParticipant) GetDataTrackCalls(stub func() types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = stub +} + +func (fake *FakeLocalParticipant) GetDataTrackReturns(result1 types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = nil + fake.getDataTrackReturns = struct { + result1 types.DataTrack + }{result1} +} + +func (fake *FakeLocalParticipant) GetDataTrackReturnsOnCall(i int, result1 types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = nil + if fake.getDataTrackReturnsOnCall == nil { + fake.getDataTrackReturnsOnCall = make(map[int]struct { + result1 types.DataTrack + }) + } + fake.getDataTrackReturnsOnCall[i] = struct { + result1 types.DataTrack + }{result1} +} + func (fake *FakeLocalParticipant) GetLogger() logger.Logger { fake.getLoggerMutex.Lock() ret, specificReturn := fake.getLoggerReturnsOnCall[len(fake.getLoggerArgsForCall)] @@ -2143,35 +2206,35 @@ func (fake *FakeLocalParticipant) OnCloseArgsForCall(i int) func(types.LocalPart return argsForCall.arg1 } -func (fake *FakeLocalParticipant) OnDataPacket(arg1 func(types.LocalParticipant, *livekit.DataPacket)) { - fake.onDataPacketMutex.Lock() - fake.onDataPacketArgsForCall = append(fake.onDataPacketArgsForCall, struct { - arg1 func(types.LocalParticipant, *livekit.DataPacket) +func (fake *FakeLocalParticipant) OnDataTrackPublished(arg1 func(types.LocalParticipant, types.DataTrack)) { + fake.onDataTrackPublishedMutex.Lock() + fake.onDataTrackPublishedArgsForCall = append(fake.onDataTrackPublishedArgsForCall, struct { + arg1 func(types.LocalParticipant, types.DataTrack) }{arg1}) - stub := fake.OnDataPacketStub - fake.recordInvocation("OnDataPacket", []interface{}{arg1}) - fake.onDataPacketMutex.Unlock() + stub := fake.OnDataTrackPublishedStub + fake.recordInvocation("OnDataTrackPublished", []interface{}{arg1}) + fake.onDataTrackPublishedMutex.Unlock() if stub != nil { - fake.OnDataPacketStub(arg1) + fake.OnDataTrackPublishedStub(arg1) } } -func (fake *FakeLocalParticipant) OnDataPacketCallCount() int { - fake.onDataPacketMutex.RLock() - defer fake.onDataPacketMutex.RUnlock() - return len(fake.onDataPacketArgsForCall) +func (fake *FakeLocalParticipant) OnDataTrackPublishedCallCount() int { + fake.onDataTrackPublishedMutex.RLock() + defer fake.onDataTrackPublishedMutex.RUnlock() + return len(fake.onDataTrackPublishedArgsForCall) } -func (fake *FakeLocalParticipant) OnDataPacketCalls(stub func(func(types.LocalParticipant, *livekit.DataPacket))) { - fake.onDataPacketMutex.Lock() - defer fake.onDataPacketMutex.Unlock() - fake.OnDataPacketStub = stub +func (fake *FakeLocalParticipant) OnDataTrackPublishedCalls(stub func(func(types.LocalParticipant, types.DataTrack))) { + fake.onDataTrackPublishedMutex.Lock() + defer fake.onDataTrackPublishedMutex.Unlock() + fake.OnDataTrackPublishedStub = stub } -func (fake *FakeLocalParticipant) OnDataPacketArgsForCall(i int) func(types.LocalParticipant, *livekit.DataPacket) { - fake.onDataPacketMutex.RLock() - defer fake.onDataPacketMutex.RUnlock() - argsForCall := fake.onDataPacketArgsForCall[i] +func (fake *FakeLocalParticipant) OnDataTrackPublishedArgsForCall(i int) func(types.LocalParticipant, types.DataTrack) { + fake.onDataTrackPublishedMutex.RLock() + defer fake.onDataTrackPublishedMutex.RUnlock() + argsForCall := fake.onDataTrackPublishedArgsForCall[i] return argsForCall.arg1 } @@ -3856,6 +3919,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.getAudioLevelMutex.RUnlock() fake.getConnectionQualityMutex.RLock() defer fake.getConnectionQualityMutex.RUnlock() + fake.getDataTrackMutex.RLock() + defer fake.getDataTrackMutex.RUnlock() fake.getLoggerMutex.RLock() defer fake.getLoggerMutex.RUnlock() fake.getPublishedTrackMutex.RLock() @@ -3890,8 +3955,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.onClaimsChangedMutex.RUnlock() fake.onCloseMutex.RLock() defer fake.onCloseMutex.RUnlock() - fake.onDataPacketMutex.RLock() - defer fake.onDataPacketMutex.RUnlock() + fake.onDataTrackPublishedMutex.RLock() + defer fake.onDataTrackPublishedMutex.RUnlock() fake.onMetadataUpdateMutex.RLock() defer fake.onMetadataUpdateMutex.RUnlock() fake.onStateChangeMutex.RLock() diff --git a/pkg/rtc/types/typesfakes/fake_participant.go b/pkg/rtc/types/typesfakes/fake_participant.go index dbba2c0a6..ebb4033cc 100644 --- a/pkg/rtc/types/typesfakes/fake_participant.go +++ b/pkg/rtc/types/typesfakes/fake_participant.go @@ -44,6 +44,16 @@ type FakeParticipant struct { debugInfoReturnsOnCall map[int]struct { result1 map[string]interface{} } + GetDataTrackStub func() types.DataTrack + getDataTrackMutex sync.RWMutex + getDataTrackArgsForCall []struct { + } + getDataTrackReturns struct { + result1 types.DataTrack + } + getDataTrackReturnsOnCall map[int]struct { + result1 types.DataTrack + } GetPublishedTrackStub func(livekit.TrackID) types.MediaTrack getPublishedTrackMutex sync.RWMutex getPublishedTrackArgsForCall []struct { @@ -373,6 +383,59 @@ func (fake *FakeParticipant) DebugInfoReturnsOnCall(i int, result1 map[string]in }{result1} } +func (fake *FakeParticipant) GetDataTrack() types.DataTrack { + fake.getDataTrackMutex.Lock() + ret, specificReturn := fake.getDataTrackReturnsOnCall[len(fake.getDataTrackArgsForCall)] + fake.getDataTrackArgsForCall = append(fake.getDataTrackArgsForCall, struct { + }{}) + stub := fake.GetDataTrackStub + fakeReturns := fake.getDataTrackReturns + fake.recordInvocation("GetDataTrack", []interface{}{}) + fake.getDataTrackMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeParticipant) GetDataTrackCallCount() int { + fake.getDataTrackMutex.RLock() + defer fake.getDataTrackMutex.RUnlock() + return len(fake.getDataTrackArgsForCall) +} + +func (fake *FakeParticipant) GetDataTrackCalls(stub func() types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = stub +} + +func (fake *FakeParticipant) GetDataTrackReturns(result1 types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = nil + fake.getDataTrackReturns = struct { + result1 types.DataTrack + }{result1} +} + +func (fake *FakeParticipant) GetDataTrackReturnsOnCall(i int, result1 types.DataTrack) { + fake.getDataTrackMutex.Lock() + defer fake.getDataTrackMutex.Unlock() + fake.GetDataTrackStub = nil + if fake.getDataTrackReturnsOnCall == nil { + fake.getDataTrackReturnsOnCall = make(map[int]struct { + result1 types.DataTrack + }) + } + fake.getDataTrackReturnsOnCall[i] = struct { + result1 types.DataTrack + }{result1} +} + func (fake *FakeParticipant) GetPublishedTrack(arg1 livekit.TrackID) types.MediaTrack { fake.getPublishedTrackMutex.Lock() ret, specificReturn := fake.getPublishedTrackReturnsOnCall[len(fake.getPublishedTrackArgsForCall)] @@ -1153,6 +1216,8 @@ func (fake *FakeParticipant) Invocations() map[string][][]interface{} { defer fake.closeMutex.RUnlock() fake.debugInfoMutex.RLock() defer fake.debugInfoMutex.RUnlock() + fake.getDataTrackMutex.RLock() + defer fake.getDataTrackMutex.RUnlock() fake.getPublishedTrackMutex.RLock() defer fake.getPublishedTrackMutex.RUnlock() fake.getPublishedTracksMutex.RLock() diff --git a/pkg/sfu/receiver.go b/pkg/sfu/receiver.go index ba6232db7..5ac5d647f 100644 --- a/pkg/sfu/receiver.go +++ b/pkg/sfu/receiver.go @@ -1,6 +1,7 @@ package sfu import ( + "errors" "io" "runtime" "sync" @@ -19,6 +20,11 @@ import ( "github.com/livekit/livekit-server/pkg/sfu/connectionquality" ) +var ( + ErrReceiverClosed = errors.New("receiver closed") + ErrDownTrackAlreadyExist = errors.New("DownTrack already exist") +) + type AudioLevelHandle func(level uint8, duration uint32) type Bitrates [DefaultMaxLayerSpatial + 1][DefaultMaxLayerTemporal + 1]int64 @@ -37,7 +43,7 @@ type TrackReceiver interface { SetUpTrackPaused(paused bool) SetMaxExpectedSpatialLayer(layer int32) - AddDownTrack(track TrackSender) + AddDownTrack(track TrackSender) error DeleteDownTrack(peerID livekit.ParticipantID) DebugInfo() map[string]interface{} @@ -285,16 +291,16 @@ func (w *WebRTCReceiver) SetUpTrackPaused(paused bool) { w.streamTrackerManager.SetPaused(paused) } -func (w *WebRTCReceiver) AddDownTrack(track TrackSender) { +func (w *WebRTCReceiver) AddDownTrack(track TrackSender) error { if w.closed.Load() { - return + return ErrReceiverClosed } w.downTrackMu.RLock() _, ok := w.index[track.PeerID()] w.downTrackMu.RUnlock() if ok { - return + return ErrDownTrackAlreadyExist } if w.Kind() == webrtc.RTPCodecTypeVideo { @@ -306,6 +312,7 @@ func (w *WebRTCReceiver) AddDownTrack(track TrackSender) { } w.storeDownTrack(track) + return nil } func (w *WebRTCReceiver) SetMaxExpectedSpatialLayer(layer int32) { diff --git a/test/multinode_roomservice_test.go b/test/multinode_roomservice_test.go index 85b9139b6..2c5dd604d 100644 --- a/test/multinode_roomservice_test.go +++ b/test/multinode_roomservice_test.go @@ -140,10 +140,10 @@ func TestMultiNodeMutePublishedTrack(t *testing.T) { Identity: identity, }) require.NoError(t, err) - if len(res.Tracks) == 2 { + if len(res.Tracks) == 3 { return "" } else { - return fmt.Sprintf("expected two tracks to be published, actual: %d", len(res.Tracks)) + return fmt.Sprintf("expected three tracks to be published, actual: %d", len(res.Tracks)) } }) diff --git a/test/singlenode_test.go b/test/singlenode_test.go index 17dae069d..ebcbf6808 100644 --- a/test/singlenode_test.go +++ b/test/singlenode_test.go @@ -420,7 +420,7 @@ func TestSingleNodeUpdateSubscriptionPermissions(t *testing.T) { if pubRemote == nil { return "could not find remote publisher" } - if len(pubRemote.Tracks) != 2 { + if len(pubRemote.Tracks) != 3 { return "did not receive metadata for published tracks" } return ""