Re-use transceiver (via ReplaceTrack) if a down track is going to be resumed. (#785)

* WIP commit

* Clean up

* spelling mistake

* Run subscribed track onBind in a go routine

* Address comments and more safety net
This commit is contained in:
Raja Subramanian
2022-06-24 15:07:48 +05:30
committed by GitHub
parent 0b630e15b6
commit adf2d191b0
9 changed files with 428 additions and 71 deletions
+110 -53
View File
@@ -162,72 +162,110 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr *
}
subTrack := NewSubscribedTrack(SubscribedTrackParams{
PublisherID: t.params.MediaTrack.PublisherID(),
PublisherIdentity: t.params.MediaTrack.PublisherIdentity(),
SubscriberID: subscriberID,
SubscriberIdentity: sub.Identity(),
MediaTrack: t.params.MediaTrack,
DownTrack: downTrack,
AdaptiveStream: sub.GetAdaptiveStream(),
PublisherID: t.params.MediaTrack.PublisherID(),
PublisherIdentity: t.params.MediaTrack.PublisherIdentity(),
Subscriber: sub,
MediaTrack: t.params.MediaTrack,
DownTrack: downTrack,
AdaptiveStream: sub.GetAdaptiveStream(),
})
// Bind callback can happen from replaceTrack, so set it up early
downTrack.OnBind(func() {
wr.DetermineReceiver(downTrack.Codec())
if err = wr.AddDownTrack(downTrack); err != nil {
t.params.Logger.Errorw("could not add down track", err, "participant", sub.Identity(), "pID", sub.ID())
}
go subTrack.Bound()
// when down track is bound, start loop to send reports
go t.sendDownTrackBindingReports(sub)
// initialize to default layer
t.notifySubscriberMaxQuality(subscriberID, downTrack.Codec(), livekit.VideoQuality_HIGH)
subTrack.SetPublisherMuted(t.params.MediaTrack.IsMuted())
})
var transceiver *webrtc.RTPTransceiver
var sender *webrtc.RTPSender
if sub.ProtocolVersion().SupportsTransceiverReuse() {
//
// AddTrack will create a new transceiver or re-use an unused one
// if the attributes match. This prevents SDP from bloating
// because of dormant transceivers building up.
//
sender, err = sub.SubscriberPC().AddTrack(downTrack)
if err != nil {
return nil, err
}
// as there is no way to get transceiver from sender, search
for _, tr := range sub.SubscriberPC().GetTransceivers() {
if tr.Sender() == sender {
transceiver = tr
break
// try cached RTP senders for a chance to replace track
replacedTrack := false
existingTransceiver := sub.GetCachedRTPTransceiver(trackID)
if existingTransceiver != nil {
rtpSender := existingTransceiver.Sender()
if rtpSender != nil {
err := rtpSender.ReplaceTrack(downTrack)
if err == nil {
sender = rtpSender
transceiver = existingTransceiver
replacedTrack = true
}
}
if transceiver == nil {
// cannot add, no transceiver
return nil, errors.New("cannot subscribe without a transceiver in place")
}
} else {
transceiver, err = sub.SubscriberPC().AddTransceiverFromTrack(downTrack, webrtc.RTPTransceiverInit{
Direction: webrtc.RTPTransceiverDirectionSendonly,
})
if err != nil {
return nil, err
}
sender = transceiver.Sender()
if sender == nil {
// cannot add, no sender
return nil, errors.New("cannot subscribe without a sender in place")
if !replacedTrack {
// Could not re-use cached transceiver for this track.
// Stop the transceiver so that it is at least not active.
// It is not usable once stopped,
//
// Adding down track will create a new transceiver (or re-use
// an inactive existing one). In either case, a renegotiation
// will happen and that will notify remote of this stopped
// transceiver
existingTransceiver.Stop()
}
}
// if cannot replace, find an unused transceiver or add new one
if transceiver == nil {
if sub.ProtocolVersion().SupportsTransceiverReuse() {
//
// AddTrack will create a new transceiver or re-use an unused one
// if the attributes match. This prevents SDP from bloating
// because of dormant transceivers building up.
//
sender, err = sub.SubscriberPC().AddTrack(downTrack)
if err != nil {
return nil, err
}
// as there is no way to get transceiver from sender, search
for _, tr := range sub.SubscriberPC().GetTransceivers() {
if tr.Sender() == sender {
transceiver = tr
break
}
}
} else {
transceiver, err = sub.SubscriberPC().AddTransceiverFromTrack(downTrack, webrtc.RTPTransceiverInit{
Direction: webrtc.RTPTransceiverDirectionSendonly,
})
if err != nil {
return nil, err
}
sender = transceiver.Sender()
}
}
if transceiver == nil {
// cannot add, no transceiver
return nil, errors.New("cannot subscribe without a transceiver in place")
}
if sender == nil {
// cannot add, no sender
return nil, errors.New("cannot subscribe without a sender in place")
}
// wthether re-using or stopping remove transceiver from cache
// NOTE: safety net, if somehow a cached transceiver is re-used by a different track
sub.UncacheRTPTransceiver(transceiver)
sendParameters := sender.GetParameters()
downTrack.SetRTPHeaderExtensions(sendParameters.HeaderExtensions)
downTrack.SetTransceiver(transceiver)
// when out track is bound, start loop to send reports
downTrack.OnBind(func() {
wr.DetermineReceiver(downTrack.Codec())
if err = wr.AddDownTrack(downTrack); err != nil {
logger.Errorw("could not add down track", err, "participant", sub.Identity(), "pID", sub.ID())
}
go subTrack.Bound()
go t.sendDownTrackBindingReports(sub)
// initialize to default layer
t.notifySubscriberMaxQuality(subscriberID, downTrack.Codec(), livekit.VideoQuality_HIGH)
subTrack.SetPublisherMuted(t.params.MediaTrack.IsMuted())
})
downTrack.OnStatsUpdate(func(_ *sfu.DownTrack, stat *livekit.AnalyticsStat) {
t.params.Telemetry.TrackStats(livekit.StreamType_DOWNSTREAM, subscriberID, trackID, stat)
})
@@ -251,7 +289,9 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr *
// since sub will lock, run it in a goroutine to avoid deadlocks
go func() {
sub.AddSubscribedTrack(subTrack)
sub.Negotiate(false)
if !replacedTrack {
sub.Negotiate(false)
}
}()
t.params.Telemetry.TrackSubscribed(context.Background(), subscriberID, t.params.MediaTrack.ToProto(),
@@ -272,7 +312,7 @@ func (t *MediaTrackSubscriptions) RemoveSubscriber(participantID livekit.Partici
t.subscribedTracksMu.Unlock()
if subTrack != nil {
subTrack.DownTrack().CloseWithFlush(!willBeResumed)
t.closeSubscribedTrack(subTrack, willBeResumed)
}
}
@@ -289,10 +329,27 @@ func (t *MediaTrackSubscriptions) RemoveAllSubscribers(willBeResumed bool) {
t.subscribedTracksMu.Unlock()
for _, subTrack := range subscribedTracks {
subTrack.DownTrack().CloseWithFlush(!willBeResumed)
t.closeSubscribedTrack(subTrack, willBeResumed)
}
}
func (t *MediaTrackSubscriptions) closeSubscribedTrack(subTrack types.SubscribedTrack, willBeResumed bool) {
dt := subTrack.DownTrack()
if dt == nil {
return
}
if willBeResumed {
tr := dt.GetTransceiver()
if tr != nil {
sub := subTrack.Subscriber()
sub.CacheRTPTransceiver(subTrack.ID(), tr)
}
}
dt.CloseWithFlush(!willBeResumed)
}
func (t *MediaTrackSubscriptions) ResyncAllSubscribers() {
t.params.Logger.Debugw("resyncing all subscribers")
+30
View File
@@ -137,6 +137,8 @@ type ParticipantImpl struct {
activeCounter atomic.Int32
firstConnected atomic.Bool
iceConfig types.IceConfig
cachedRTPTransceivers map[livekit.TrackID]*webrtc.RTPTransceiver
}
func NewParticipant(params ParticipantParams) (*ParticipantImpl, error) {
@@ -159,6 +161,7 @@ func NewParticipant(params ParticipantParams) (*ParticipantImpl, error) {
subscribedTo: make(map[livekit.ParticipantID]struct{}),
connectedAt: time.Now(),
rttUpdatedAt: time.Now(),
cachedRTPTransceivers: make(map[livekit.TrackID]*webrtc.RTPTransceiver),
}
p.version.Store(params.InitialVersion)
p.migrateState.Store(types.MigrateStateInit)
@@ -1908,3 +1911,30 @@ func (p *ParticipantImpl) setDowntracksConnected() {
}
}
}
func (p *ParticipantImpl) CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver) {
p.lock.Lock()
if existing := p.cachedRTPTransceivers[trackID]; existing != nil && existing != rtpTransceiver {
p.params.Logger.Infow("cached transceiver change", "trackID", trackID)
}
p.cachedRTPTransceivers[trackID] = rtpTransceiver
p.lock.Unlock()
}
func (p *ParticipantImpl) UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver) {
p.lock.Lock()
for trackID, tr := range p.cachedRTPTransceivers {
if tr == rtpTransceiver {
delete(p.cachedRTPTransceivers, trackID)
break
}
}
p.lock.Unlock()
}
func (p *ParticipantImpl) GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver {
p.lock.RLock()
defer p.lock.RUnlock()
return p.cachedRTPTransceivers[trackID]
}
+24 -13
View File
@@ -19,13 +19,12 @@ const (
)
type SubscribedTrackParams struct {
PublisherID livekit.ParticipantID
PublisherIdentity livekit.ParticipantIdentity
SubscriberID livekit.ParticipantID
SubscriberIdentity livekit.ParticipantIdentity
MediaTrack types.MediaTrack
DownTrack *sfu.DownTrack
AdaptiveStream bool
PublisherID livekit.ParticipantID
PublisherIdentity livekit.ParticipantIdentity
Subscriber types.LocalParticipant
MediaTrack types.MediaTrack
DownTrack *sfu.DownTrack
AdaptiveStream bool
}
type SubscribedTrack struct {
@@ -34,7 +33,8 @@ type SubscribedTrack struct {
pubMuted atomic.Bool
settings atomic.Value // *livekit.UpdateTrackSettings
onBind func()
onBind atomic.Value // func()
bound atomic.Bool
debouncer func(func())
}
@@ -49,15 +49,22 @@ func NewSubscribedTrack(params SubscribedTrackParams) *SubscribedTrack {
}
func (t *SubscribedTrack) OnBind(f func()) {
t.onBind = f
t.onBind.Store(f)
t.maybeOnBind()
}
func (t *SubscribedTrack) Bound() {
t.bound.Store(true)
if !t.params.AdaptiveStream {
t.params.DownTrack.SetMaxSpatialLayer(utils.SpatialLayerForQuality(livekit.VideoQuality_HIGH))
}
if t.onBind != nil {
t.onBind()
t.maybeOnBind()
}
func (t *SubscribedTrack) maybeOnBind() {
if onBind := t.onBind.Load(); onBind != nil && t.bound.Load() {
go onBind.(func())()
}
}
@@ -74,11 +81,15 @@ func (t *SubscribedTrack) PublisherIdentity() livekit.ParticipantIdentity {
}
func (t *SubscribedTrack) SubscriberID() livekit.ParticipantID {
return t.params.SubscriberID
return t.params.Subscriber.ID()
}
func (t *SubscribedTrack) SubscriberIdentity() livekit.ParticipantIdentity {
return t.params.SubscriberIdentity
return t.params.Subscriber.Identity()
}
func (t *SubscribedTrack) Subscriber() types.LocalParticipant {
return t.params.Subscriber
}
func (t *SubscribedTrack) DownTrack() *sfu.DownTrack {
+5
View File
@@ -246,6 +246,10 @@ type LocalParticipant interface {
SetMigrateInfo(previousAnswer *webrtc.SessionDescription, mediaTracks []*livekit.TrackPublishedResponse, dataChannels []*livekit.DataChannelInfo)
UpdateRTT(rtt uint32)
CacheRTPTransceiver(trackID livekit.TrackID, rtpTransceiver *webrtc.RTPTransceiver)
UncacheRTPTransceiver(rtpTransceiver *webrtc.RTPTransceiver)
GetCachedRTPTransceiver(trackID livekit.TrackID) *webrtc.RTPTransceiver
}
// Room is a container of participants, and can provide room-level actions
@@ -325,6 +329,7 @@ type SubscribedTrack interface {
PublisherIdentity() livekit.ParticipantIdentity
SubscriberID() livekit.ParticipantID
SubscriberIdentity() livekit.ParticipantIdentity
Subscriber() LocalParticipant
DownTrack() *sfu.DownTrack
MediaTrack() MediaTrack
IsMuted() bool
@@ -50,6 +50,12 @@ type FakeLocalParticipant struct {
addTrackArgsForCall []struct {
arg1 *livekit.AddTrackRequest
}
CacheRTPTransceiverStub func(livekit.TrackID, *webrtc.RTPTransceiver)
cacheRTPTransceiverMutex sync.RWMutex
cacheRTPTransceiverArgsForCall []struct {
arg1 livekit.TrackID
arg2 *webrtc.RTPTransceiver
}
CanPublishStub func() bool
canPublishMutex sync.RWMutex
canPublishArgsForCall []struct {
@@ -144,6 +150,17 @@ type FakeLocalParticipant struct {
result1 float64
result2 bool
}
GetCachedRTPTransceiverStub func(livekit.TrackID) *webrtc.RTPTransceiver
getCachedRTPTransceiverMutex sync.RWMutex
getCachedRTPTransceiverArgsForCall []struct {
arg1 livekit.TrackID
}
getCachedRTPTransceiverReturns struct {
result1 *webrtc.RTPTransceiver
}
getCachedRTPTransceiverReturnsOnCall map[int]struct {
result1 *webrtc.RTPTransceiver
}
GetConnectionQualityStub func() *livekit.ConnectionQualityInfo
getConnectionQualityMutex sync.RWMutex
getConnectionQualityArgsForCall []struct {
@@ -589,6 +606,11 @@ type FakeLocalParticipant struct {
toProtoReturnsOnCall map[int]struct {
result1 *livekit.ParticipantInfo
}
UncacheRTPTransceiverStub func(*webrtc.RTPTransceiver)
uncacheRTPTransceiverMutex sync.RWMutex
uncacheRTPTransceiverArgsForCall []struct {
arg1 *webrtc.RTPTransceiver
}
UpdateMediaLossStub func(livekit.NodeID, livekit.TrackID, uint32) error
updateMediaLossMutex sync.RWMutex
updateMediaLossArgsForCall []struct {
@@ -851,6 +873,39 @@ func (fake *FakeLocalParticipant) AddTrackArgsForCall(i int) *livekit.AddTrackRe
return argsForCall.arg1
}
func (fake *FakeLocalParticipant) CacheRTPTransceiver(arg1 livekit.TrackID, arg2 *webrtc.RTPTransceiver) {
fake.cacheRTPTransceiverMutex.Lock()
fake.cacheRTPTransceiverArgsForCall = append(fake.cacheRTPTransceiverArgsForCall, struct {
arg1 livekit.TrackID
arg2 *webrtc.RTPTransceiver
}{arg1, arg2})
stub := fake.CacheRTPTransceiverStub
fake.recordInvocation("CacheRTPTransceiver", []interface{}{arg1, arg2})
fake.cacheRTPTransceiverMutex.Unlock()
if stub != nil {
fake.CacheRTPTransceiverStub(arg1, arg2)
}
}
func (fake *FakeLocalParticipant) CacheRTPTransceiverCallCount() int {
fake.cacheRTPTransceiverMutex.RLock()
defer fake.cacheRTPTransceiverMutex.RUnlock()
return len(fake.cacheRTPTransceiverArgsForCall)
}
func (fake *FakeLocalParticipant) CacheRTPTransceiverCalls(stub func(livekit.TrackID, *webrtc.RTPTransceiver)) {
fake.cacheRTPTransceiverMutex.Lock()
defer fake.cacheRTPTransceiverMutex.Unlock()
fake.CacheRTPTransceiverStub = stub
}
func (fake *FakeLocalParticipant) CacheRTPTransceiverArgsForCall(i int) (livekit.TrackID, *webrtc.RTPTransceiver) {
fake.cacheRTPTransceiverMutex.RLock()
defer fake.cacheRTPTransceiverMutex.RUnlock()
argsForCall := fake.cacheRTPTransceiverArgsForCall[i]
return argsForCall.arg1, argsForCall.arg2
}
func (fake *FakeLocalParticipant) CanPublish() bool {
fake.canPublishMutex.Lock()
ret, specificReturn := fake.canPublishReturnsOnCall[len(fake.canPublishArgsForCall)]
@@ -1340,6 +1395,67 @@ func (fake *FakeLocalParticipant) GetAudioLevelReturnsOnCall(i int, result1 floa
}{result1, result2}
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiver(arg1 livekit.TrackID) *webrtc.RTPTransceiver {
fake.getCachedRTPTransceiverMutex.Lock()
ret, specificReturn := fake.getCachedRTPTransceiverReturnsOnCall[len(fake.getCachedRTPTransceiverArgsForCall)]
fake.getCachedRTPTransceiverArgsForCall = append(fake.getCachedRTPTransceiverArgsForCall, struct {
arg1 livekit.TrackID
}{arg1})
stub := fake.GetCachedRTPTransceiverStub
fakeReturns := fake.getCachedRTPTransceiverReturns
fake.recordInvocation("GetCachedRTPTransceiver", []interface{}{arg1})
fake.getCachedRTPTransceiverMutex.Unlock()
if stub != nil {
return stub(arg1)
}
if specificReturn {
return ret.result1
}
return fakeReturns.result1
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCallCount() int {
fake.getCachedRTPTransceiverMutex.RLock()
defer fake.getCachedRTPTransceiverMutex.RUnlock()
return len(fake.getCachedRTPTransceiverArgsForCall)
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiverCalls(stub func(livekit.TrackID) *webrtc.RTPTransceiver) {
fake.getCachedRTPTransceiverMutex.Lock()
defer fake.getCachedRTPTransceiverMutex.Unlock()
fake.GetCachedRTPTransceiverStub = stub
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiverArgsForCall(i int) livekit.TrackID {
fake.getCachedRTPTransceiverMutex.RLock()
defer fake.getCachedRTPTransceiverMutex.RUnlock()
argsForCall := fake.getCachedRTPTransceiverArgsForCall[i]
return argsForCall.arg1
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiverReturns(result1 *webrtc.RTPTransceiver) {
fake.getCachedRTPTransceiverMutex.Lock()
defer fake.getCachedRTPTransceiverMutex.Unlock()
fake.GetCachedRTPTransceiverStub = nil
fake.getCachedRTPTransceiverReturns = struct {
result1 *webrtc.RTPTransceiver
}{result1}
}
func (fake *FakeLocalParticipant) GetCachedRTPTransceiverReturnsOnCall(i int, result1 *webrtc.RTPTransceiver) {
fake.getCachedRTPTransceiverMutex.Lock()
defer fake.getCachedRTPTransceiverMutex.Unlock()
fake.GetCachedRTPTransceiverStub = nil
if fake.getCachedRTPTransceiverReturnsOnCall == nil {
fake.getCachedRTPTransceiverReturnsOnCall = make(map[int]struct {
result1 *webrtc.RTPTransceiver
})
}
fake.getCachedRTPTransceiverReturnsOnCall[i] = struct {
result1 *webrtc.RTPTransceiver
}{result1}
}
func (fake *FakeLocalParticipant) GetConnectionQuality() *livekit.ConnectionQualityInfo {
fake.getConnectionQualityMutex.Lock()
ret, specificReturn := fake.getConnectionQualityReturnsOnCall[len(fake.getConnectionQualityArgsForCall)]
@@ -3805,6 +3921,38 @@ func (fake *FakeLocalParticipant) ToProtoReturnsOnCall(i int, result1 *livekit.P
}{result1}
}
func (fake *FakeLocalParticipant) UncacheRTPTransceiver(arg1 *webrtc.RTPTransceiver) {
fake.uncacheRTPTransceiverMutex.Lock()
fake.uncacheRTPTransceiverArgsForCall = append(fake.uncacheRTPTransceiverArgsForCall, struct {
arg1 *webrtc.RTPTransceiver
}{arg1})
stub := fake.UncacheRTPTransceiverStub
fake.recordInvocation("UncacheRTPTransceiver", []interface{}{arg1})
fake.uncacheRTPTransceiverMutex.Unlock()
if stub != nil {
fake.UncacheRTPTransceiverStub(arg1)
}
}
func (fake *FakeLocalParticipant) UncacheRTPTransceiverCallCount() int {
fake.uncacheRTPTransceiverMutex.RLock()
defer fake.uncacheRTPTransceiverMutex.RUnlock()
return len(fake.uncacheRTPTransceiverArgsForCall)
}
func (fake *FakeLocalParticipant) UncacheRTPTransceiverCalls(stub func(*webrtc.RTPTransceiver)) {
fake.uncacheRTPTransceiverMutex.Lock()
defer fake.uncacheRTPTransceiverMutex.Unlock()
fake.UncacheRTPTransceiverStub = stub
}
func (fake *FakeLocalParticipant) UncacheRTPTransceiverArgsForCall(i int) *webrtc.RTPTransceiver {
fake.uncacheRTPTransceiverMutex.RLock()
defer fake.uncacheRTPTransceiverMutex.RUnlock()
argsForCall := fake.uncacheRTPTransceiverArgsForCall[i]
return argsForCall.arg1
}
func (fake *FakeLocalParticipant) UpdateMediaLoss(arg1 livekit.NodeID, arg2 livekit.TrackID, arg3 uint32) error {
fake.updateMediaLossMutex.Lock()
ret, specificReturn := fake.updateMediaLossReturnsOnCall[len(fake.updateMediaLossArgsForCall)]
@@ -4165,6 +4313,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} {
defer fake.addSubscriberMutex.RUnlock()
fake.addTrackMutex.RLock()
defer fake.addTrackMutex.RUnlock()
fake.cacheRTPTransceiverMutex.RLock()
defer fake.cacheRTPTransceiverMutex.RUnlock()
fake.canPublishMutex.RLock()
defer fake.canPublishMutex.RUnlock()
fake.canPublishDataMutex.RLock()
@@ -4183,6 +4333,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} {
defer fake.getAdaptiveStreamMutex.RUnlock()
fake.getAudioLevelMutex.RLock()
defer fake.getAudioLevelMutex.RUnlock()
fake.getCachedRTPTransceiverMutex.RLock()
defer fake.getCachedRTPTransceiverMutex.RUnlock()
fake.getConnectionQualityMutex.RLock()
defer fake.getConnectionQualityMutex.RUnlock()
fake.getLoggerMutex.RLock()
@@ -4285,6 +4437,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} {
defer fake.subscriptionPermissionUpdateMutex.RUnlock()
fake.toProtoMutex.RLock()
defer fake.toProtoMutex.RUnlock()
fake.uncacheRTPTransceiverMutex.RLock()
defer fake.uncacheRTPTransceiverMutex.RUnlock()
fake.updateMediaLossMutex.RLock()
defer fake.updateMediaLossMutex.RUnlock()
fake.updateRTTMutex.RLock()
@@ -80,6 +80,16 @@ type FakeSubscribedTrack struct {
setPublisherMutedArgsForCall []struct {
arg1 bool
}
SubscriberStub func() types.LocalParticipant
subscriberMutex sync.RWMutex
subscriberArgsForCall []struct {
}
subscriberReturns struct {
result1 types.LocalParticipant
}
subscriberReturnsOnCall map[int]struct {
result1 types.LocalParticipant
}
SubscriberIDStub func() livekit.ParticipantID
subscriberIDMutex sync.RWMutex
subscriberIDArgsForCall []struct {
@@ -495,6 +505,59 @@ func (fake *FakeSubscribedTrack) SetPublisherMutedArgsForCall(i int) bool {
return argsForCall.arg1
}
func (fake *FakeSubscribedTrack) Subscriber() types.LocalParticipant {
fake.subscriberMutex.Lock()
ret, specificReturn := fake.subscriberReturnsOnCall[len(fake.subscriberArgsForCall)]
fake.subscriberArgsForCall = append(fake.subscriberArgsForCall, struct {
}{})
stub := fake.SubscriberStub
fakeReturns := fake.subscriberReturns
fake.recordInvocation("Subscriber", []interface{}{})
fake.subscriberMutex.Unlock()
if stub != nil {
return stub()
}
if specificReturn {
return ret.result1
}
return fakeReturns.result1
}
func (fake *FakeSubscribedTrack) SubscriberCallCount() int {
fake.subscriberMutex.RLock()
defer fake.subscriberMutex.RUnlock()
return len(fake.subscriberArgsForCall)
}
func (fake *FakeSubscribedTrack) SubscriberCalls(stub func() types.LocalParticipant) {
fake.subscriberMutex.Lock()
defer fake.subscriberMutex.Unlock()
fake.SubscriberStub = stub
}
func (fake *FakeSubscribedTrack) SubscriberReturns(result1 types.LocalParticipant) {
fake.subscriberMutex.Lock()
defer fake.subscriberMutex.Unlock()
fake.SubscriberStub = nil
fake.subscriberReturns = struct {
result1 types.LocalParticipant
}{result1}
}
func (fake *FakeSubscribedTrack) SubscriberReturnsOnCall(i int, result1 types.LocalParticipant) {
fake.subscriberMutex.Lock()
defer fake.subscriberMutex.Unlock()
fake.SubscriberStub = nil
if fake.subscriberReturnsOnCall == nil {
fake.subscriberReturnsOnCall = make(map[int]struct {
result1 types.LocalParticipant
})
}
fake.subscriberReturnsOnCall[i] = struct {
result1 types.LocalParticipant
}{result1}
}
func (fake *FakeSubscribedTrack) SubscriberID() livekit.ParticipantID {
fake.subscriberIDMutex.Lock()
ret, specificReturn := fake.subscriberIDReturnsOnCall[len(fake.subscriberIDArgsForCall)]
@@ -676,6 +739,8 @@ func (fake *FakeSubscribedTrack) Invocations() map[string][][]interface{} {
defer fake.publisherIdentityMutex.RUnlock()
fake.setPublisherMutedMutex.RLock()
defer fake.setPublisherMutedMutex.RUnlock()
fake.subscriberMutex.RLock()
defer fake.subscriberMutex.RUnlock()
fake.subscriberIDMutex.RLock()
defer fake.subscriberIDMutex.RUnlock()
fake.subscriberIdentityMutex.RLock()
+33 -4
View File
@@ -67,8 +67,15 @@ type DummyReceiver struct {
streamId string
codec webrtc.RTPCodecParameters
headerExtensions []webrtc.RTPHeaderExtensionParameter
downtrackLock sync.Mutex
downtracks map[livekit.ParticipantID]sfu.TrackSender
downtrackLock sync.Mutex
downtracks map[livekit.ParticipantID]sfu.TrackSender
settingsLock sync.Mutex
maxExpectedLayerValid bool
maxExpectedLayer int32
pausedValid bool
paused bool
}
func NewDummyReceiver(trackID livekit.TrackID, streamId string, codec webrtc.RTPCodecParameters, headerExtensions []webrtc.RTPHeaderExtensionParameter) *DummyReceiver {
@@ -87,13 +94,25 @@ func (d *DummyReceiver) Receiver() sfu.TrackReceiver {
}
func (d *DummyReceiver) Upgrade(receiver sfu.TrackReceiver) {
d.downtrackLock.Lock()
defer d.downtrackLock.Unlock()
d.receiver.CompareAndSwap(nil, receiver)
d.downtrackLock.Lock()
for _, t := range d.downtracks {
receiver.AddDownTrack(t)
}
d.downtracks = make(map[livekit.ParticipantID]sfu.TrackSender)
d.downtrackLock.Unlock()
d.settingsLock.Lock()
if d.maxExpectedLayerValid {
receiver.SetMaxExpectedSpatialLayer(d.maxExpectedLayer)
}
d.maxExpectedLayerValid = false
if d.pausedValid {
receiver.SetUpTrackPaused(d.paused)
}
d.pausedValid = false
d.settingsLock.Unlock()
}
func (d *DummyReceiver) TrackID() livekit.TrackID {
@@ -148,12 +167,22 @@ func (d *DummyReceiver) SendPLI(layer int32, force bool) {
func (d *DummyReceiver) SetUpTrackPaused(paused bool) {
if r, ok := d.receiver.Load().(sfu.TrackReceiver); ok {
r.SetUpTrackPaused(paused)
} else {
d.settingsLock.Lock()
d.pausedValid = true
d.paused = paused
d.settingsLock.Unlock()
}
}
func (d *DummyReceiver) SetMaxExpectedSpatialLayer(layer int32) {
if r, ok := d.receiver.Load().(sfu.TrackReceiver); ok {
r.SetMaxExpectedSpatialLayer(layer)
} else {
d.settingsLock.Lock()
d.maxExpectedLayerValid = true
d.maxExpectedLayer = layer
d.settingsLock.Unlock()
}
}
+4
View File
@@ -364,6 +364,10 @@ func (d *DownTrack) SetTransceiver(transceiver *webrtc.RTPTransceiver) {
d.transceiver = transceiver
}
func (d *DownTrack) GetTransceiver() *webrtc.RTPTransceiver {
return d.transceiver
}
func (d *DownTrack) maybeStartKeyFrameRequester() {
//
// Always move to next generation to abandon any running key frame requester
+3 -1
View File
@@ -240,7 +240,9 @@ func (s *StreamAllocator) AddTrack(downTrack *DownTrack, params AddTrackParams)
func (s *StreamAllocator) RemoveTrack(downTrack *DownTrack) {
s.videoTracksMu.Lock()
delete(s.videoTracks, livekit.TrackID(downTrack.ID()))
if existing := s.videoTracks[livekit.TrackID(downTrack.ID())]; existing != nil && existing.DownTrack() == downTrack {
delete(s.videoTracks, livekit.TrackID(downTrack.ID()))
}
s.videoTracksMu.Unlock()
// LK-TODO: use any saved bandwidth to re-distribute