Datatrack for data channel (#476)

* data track
This commit is contained in:
cnderrauber
2022-03-01 15:48:20 +08:00
committed by GitHub
parent 13e21e7c45
commit 913ef3a646
11 changed files with 738 additions and 88 deletions
+189
View File
@@ -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{}{}
}
+40 -49
View File
@@ -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
+6 -2
View File
@@ -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)
+9 -3
View File
@@ -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 {
+12 -1
View File
@@ -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
}
+312
View File
@@ -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)
@@ -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()
@@ -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()
+11 -4
View File
@@ -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) {
+2 -2
View File
@@ -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))
}
})
+1 -1
View File
@@ -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 ""