Handle repair SSRC of simulcast tracks during migration. (#4193)

* Handle repair SSRC of simulcast tracks during migration.

* fix

* fix comment
This commit is contained in:
Raja Subramanian
2025-12-25 14:45:48 +05:30
committed by GitHub
parent c6bf7a2786
commit ed8e6afcd7
13 changed files with 344 additions and 180 deletions
+1 -1
View File
@@ -23,7 +23,7 @@ require (
github.com/jxskiss/base62 v1.1.0
github.com/livekit/mageutil v0.0.0-20250511045019-0f1ff63f7731
github.com/livekit/mediatransportutil v0.0.0-20251213100503-cc390ae365e9
github.com/livekit/protocol v1.43.5-0.20251222225221-fa169ac100d9
github.com/livekit/protocol v1.43.5-0.20251225082431-862d938ee998
github.com/livekit/psrpc v0.7.1
github.com/mackerelio/go-osstat v0.2.6
github.com/magefile/mage v1.15.0
+2 -2
View File
@@ -173,8 +173,8 @@ github.com/livekit/mageutil v0.0.0-20250511045019-0f1ff63f7731 h1:9x+U2HGLrSw5AT
github.com/livekit/mageutil v0.0.0-20250511045019-0f1ff63f7731/go.mod h1:Rs3MhFwutWhGwmY1VQsygw28z5bWcnEYmS1OG9OxjOQ=
github.com/livekit/mediatransportutil v0.0.0-20251213100503-cc390ae365e9 h1:ciqzzn+oEex3mCa1n1GmlQrv+ZkGpgUbQPSG3PD0htM=
github.com/livekit/mediatransportutil v0.0.0-20251213100503-cc390ae365e9/go.mod h1:mSNtYzSf6iY9xM3UX42VEI+STHvMgHmrYzEHPcdhB8A=
github.com/livekit/protocol v1.43.5-0.20251222225221-fa169ac100d9 h1:SHaCOwMj3MPOeouW+hSIN/UkL6LAqwQOt6WD53NKuEQ=
github.com/livekit/protocol v1.43.5-0.20251222225221-fa169ac100d9/go.mod h1:n00Ul4P6o2YILGhxw+O57B0h/bF3Je9PzRN36fElCmw=
github.com/livekit/protocol v1.43.5-0.20251225082431-862d938ee998 h1:6+NrF16bkaExumkVjjf0JE9LLlPk5A4qZbSvCmdfRco=
github.com/livekit/protocol v1.43.5-0.20251225082431-862d938ee998/go.mod h1:n00Ul4P6o2YILGhxw+O57B0h/bF3Je9PzRN36fElCmw=
github.com/livekit/psrpc v0.7.1 h1:ms37az0QTD3UXIWuUC5D/SkmKOlRMVRsI261eBWu/Vw=
github.com/livekit/psrpc v0.7.1/go.mod h1:bZ4iHFQptTkbPnB0LasvRNu/OBYXEu1NA6O5BMFo9kk=
github.com/mackerelio/go-osstat v0.2.6 h1:gs4U8BZeS1tjrL08tt5VUliVvSWP26Ai2Ob8Lr7f2i0=
+7 -4
View File
@@ -34,6 +34,7 @@ import (
"github.com/livekit/livekit-server/pkg/sfu"
"github.com/livekit/livekit-server/pkg/sfu/buffer"
"github.com/livekit/livekit-server/pkg/sfu/connectionquality"
"github.com/livekit/livekit-server/pkg/sfu/interceptor"
"github.com/livekit/livekit-server/pkg/sfu/mime"
"github.com/livekit/livekit-server/pkg/telemetry"
util "github.com/livekit/mediatransportutil"
@@ -87,7 +88,7 @@ type MediaTrackParams struct {
Telemetry telemetry.TelemetryService
Logger logger.Logger
Reporter roomobs.TrackReporter
SimTracks map[uint32]SimulcastTrackInfo
SimTracks map[uint32]interceptor.SimulcastTrackInfo
OnRTCP func([]rtcp.Packet)
ForwardStats *sfu.ForwardStats
OnTrackEverSubscribed func(livekit.TrackID)
@@ -492,8 +493,8 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track sfu.TrackRe
t.MediaTrackReceiver.SetupReceiver(newWR, priority, mid)
for ssrc, info := range t.params.SimTracks {
if info.Mid == mid {
t.MediaTrackReceiver.SetLayerSsrc(mimeType, info.Rid, ssrc)
if info.Mid == mid && !info.IsRepairStream {
t.MediaTrackReceiver.SetLayerSsrcsForRid(mimeType, info.StreamID, ssrc, info.RepairSSRC)
}
}
wr = newWR
@@ -549,7 +550,7 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track sfu.TrackRe
return newCodec, false
}
t.MediaTrackReceiver.SetLayerSsrc(mimeType, track.RID(), uint32(track.SSRC()))
t.MediaTrackReceiver.SetLayerSsrcsForRid(mimeType, track.RID(), uint32(track.SSRC()), 0)
if regressCodec {
for _, c := range ti.Codecs {
@@ -566,6 +567,8 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track sfu.TrackRe
}
}
buff.OnNotifyRTX(t.MediaTrackReceiver.setLayerRtxInfo)
// if subscriber request fps before fps calculated, update them after fps updated.
buff.OnFpsChanged(func() {
t.MediaTrackSubscriptions.UpdateVideoLayers()
+97 -12
View File
@@ -18,6 +18,7 @@ import (
"context"
"errors"
"fmt"
"slices"
"sort"
"strings"
"sync"
@@ -25,7 +26,6 @@ import (
"github.com/pion/rtcp"
"github.com/pion/webrtc/v4"
"go.uber.org/atomic"
"golang.org/x/exp/slices"
"google.golang.org/protobuf/proto"
"github.com/livekit/protocol/livekit"
@@ -314,7 +314,7 @@ func (t *MediaTrackReceiver) HandleReceiverCodecChange(r sfu.TrackReceiver, code
// remove old codec from potential codecs
for i, c := range t.potentialCodecs {
if strings.EqualFold(c.MimeType, codec.MimeType) {
slices.Delete(t.potentialCodecs, i, i+1)
t.potentialCodecs = slices.Delete(t.potentialCodecs, i, i+1)
break
}
}
@@ -632,14 +632,7 @@ func (t *MediaTrackReceiver) RevokeDisallowedSubscribers(allowedSubscriberIdenti
continue
}
found := false
for _, allowedIdentity := range allowedSubscriberIdentities {
if subTrack.SubscriberIdentity() == allowedIdentity {
found = true
break
}
}
found := slices.Contains(allowedSubscriberIdentities, subTrack.SubscriberIdentity())
if !found {
t.params.Logger.Infow("revoking subscription",
"subscriber", subTrack.SubscriberIdentity(),
@@ -660,7 +653,7 @@ func (t *MediaTrackReceiver) updateTrackInfoOfReceivers() {
}
}
func (t *MediaTrackReceiver) SetLayerSsrc(mimeType mime.MimeType, rid string, ssrc uint32) {
func (t *MediaTrackReceiver) SetLayerSsrcsForRid(mimeType mime.MimeType, rid string, ssrc uint32, repairSSRC uint32) {
t.lock.Lock()
trackInfo := t.TrackInfoClone()
layer := buffer.GetSpatialLayerForRid(mimeType, rid, trackInfo)
@@ -689,6 +682,20 @@ func (t *MediaTrackReceiver) SetLayerSsrc(mimeType mime.MimeType, rid string, ss
}
if !ssrcFound && matchingLayer != nil {
matchingLayer.Ssrc = ssrc
if repairSSRC != 0 {
matchingLayer.RepairSsrc = repairSSRC
}
}
if ssrcFound {
t.params.Logger.Warnw(
"not overriding ssrc", nil,
"rid", rid,
"ssrc", ssrc,
"existingSSRC", matchingLayer.Ssrc,
"repairSSRC", repairSSRC,
"existingRepairSSRC", matchingLayer.RepairSsrc,
"trackInfo", trackInfo,
)
}
// for client don't use simulcast codecs (old client version or single codec)
@@ -703,6 +710,77 @@ func (t *MediaTrackReceiver) SetLayerSsrc(mimeType mime.MimeType, rid string, ss
t.updateTrackInfoOfReceivers()
}
func (t *MediaTrackReceiver) setLayerRtxInfo(ssrc uint32, repairSSRC uint32, rsid string) {
t.params.Logger.Debugw("rtx notification", "ssrc", ssrc, "repairSSRC", repairSSRC, "rsid", rsid)
if ssrc == 0 || repairSSRC == 0 || rsid == "" {
return
}
t.lock.Lock()
trackInfo := t.TrackInfoClone()
done:
for _, ci := range trackInfo.Codecs {
for _, l := range ci.Layers {
if l.Ssrc == ssrc {
if (l.RepairSsrc != 0 && l.RepairSsrc != repairSSRC) || (l.Rid != "" && l.Rid != rsid) {
t.params.Logger.Warnw(
"not overriding rtx info", nil,
"ssrc", ssrc,
"repairSSRC", repairSSRC,
"existingRepairSSRC", l.RepairSsrc,
"rsid", rsid,
"existingRid", l.Rid,
"trackInfo", logger.Proto(trackInfo),
)
} else {
l.RepairSsrc = repairSSRC
t.params.Logger.Debugw(
"set rtx info",
"ssrc", ssrc,
"repairSSRC", repairSSRC,
"rsid", rsid,
"trackInfo", logger.Proto(trackInfo),
)
}
break done
}
}
}
// backwards compatibility
for _, l := range trackInfo.Layers {
if l.Ssrc == ssrc {
if (l.RepairSsrc != 0 && l.RepairSsrc != repairSSRC) || (l.Rid != "" && l.Rid != rsid) {
t.params.Logger.Warnw(
"not overriding rtx info", nil,
"ssrc", ssrc,
"repairSSRC", repairSSRC,
"existingRepairSSRC", l.RepairSsrc,
"rsid", rsid,
"existingRid", l.Rid,
"trackInfo", logger.Proto(trackInfo),
)
} else {
l.RepairSsrc = repairSSRC
t.params.Logger.Debugw(
"set rtx info",
"ssrc", ssrc,
"repairSSRC", repairSSRC,
"rsid", rsid,
"trackInfo", logger.Proto(trackInfo),
)
}
break
}
}
t.trackInfo.Store(trackInfo)
t.lock.Unlock()
// change not propagated as it is internal
}
func (t *MediaTrackReceiver) UpdateCodecInfo(codecs []*livekit.SimulcastCodec) {
t.lock.Lock()
trackInfo := t.TrackInfoClone()
@@ -778,7 +856,7 @@ func (t *MediaTrackReceiver) UpdateTrackInfo(ti *livekit.TrackInfo) {
t.lock.Lock()
trackInfo := t.TrackInfo()
// patch Mid and SSRC of codecs/layers by keeping original if available
// patch Mid/Rid and Ssrc/RtxSsrc of codecs/layers by keeping original if available
for i, ci := range clonedInfo.Codecs {
for _, originCi := range trackInfo.Codecs {
if !mime.IsMimeTypeStringEqual(ci.MimeType, originCi.MimeType) {
@@ -795,6 +873,13 @@ func (t *MediaTrackReceiver) UpdateTrackInfo(ti *livekit.TrackInfo) {
if originLayer.Ssrc != 0 {
layer.Ssrc = originLayer.Ssrc
}
if originLayer.Rid != "" {
layer.Rid = originLayer.Rid
}
if originLayer.RepairSsrc != 0 {
layer.RepairSsrc = originLayer.RepairSsrc
}
break
}
}
+10 -11
View File
@@ -60,6 +60,7 @@ import (
"github.com/livekit/livekit-server/pkg/sfu"
"github.com/livekit/livekit-server/pkg/sfu/buffer"
"github.com/livekit/livekit-server/pkg/sfu/connectionquality"
"github.com/livekit/livekit-server/pkg/sfu/interceptor"
"github.com/livekit/livekit-server/pkg/sfu/mime"
"github.com/livekit/livekit-server/pkg/sfu/pacer"
"github.com/livekit/livekit-server/pkg/sfu/streamallocator"
@@ -179,7 +180,7 @@ type ParticipantParams struct {
LoggerResolver logger.DeferredFieldResolver
Reporter roomobs.ParticipantSessionReporter
ReporterResolver roomobs.ParticipantReporterResolver
SimTracks map[uint32]SimulcastTrackInfo
SimTracks map[uint32]interceptor.SimulcastTrackInfo
Grants *auth.ClaimGrants
InitialVersion uint32
ClientConf *livekit.ClientConfiguration
@@ -980,7 +981,7 @@ func (p *ParticipantImpl) synthesizeAddTrackRequests(parsedOffer *sdp.SessionDes
if ridsOk {
// add simulcast layers, NOTE: only quality can be set as dimensions/fps is not available
n := min(len(rids), int(buffer.DefaultMaxLayerSpatial)+1)
for i := 0; i < n; i++ {
for i := range n {
// WARN: casting int -> protobuf enum
req.Layers = append(req.Layers, &livekit.VideoLayer{Quality: livekit.VideoQuality(i)})
}
@@ -1023,7 +1024,7 @@ func (p *ParticipantImpl) updateRidsFromSDP(parsed *sdp.SessionDescription, unma
}
outRids = buffer.NormalizeVideoLayersRid(outRids)
} else {
for i := 0; i < len(inRids); i++ {
for i := range len(inRids) {
outRids[i] = ""
}
}
@@ -3062,12 +3063,10 @@ func (p *ParticipantImpl) setTrackMuted(mute *livekit.MuteTrackRequest, fromAdmi
return trackInfo
}
func (p *ParticipantImpl) mediaTrackReceived(track sfu.TrackRemote, rtpReceiver *webrtc.RTPReceiver) (
*MediaTrack,
bool,
bool,
buffer.VideoLayersRid,
) {
func (p *ParticipantImpl) mediaTrackReceived(
track sfu.TrackRemote,
rtpReceiver *webrtc.RTPReceiver,
) (*MediaTrack, bool, bool, buffer.VideoLayersRid) {
p.pendingTracksLock.Lock()
newTrack := false
@@ -3246,8 +3245,8 @@ func (p *ParticipantImpl) addMigratedTrack(cid string, ti *livekit.TrackInfo) *M
for _, codec := range ti.Codecs {
for ssrc, info := range p.params.SimTracks {
if info.Mid == codec.Mid {
mt.SetLayerSsrc(mime.NormalizeMimeType(codec.MimeType), info.Rid, ssrc)
if info.Mid == codec.Mid && !info.IsRepairStream {
mt.SetLayerSsrcsForRid(mime.NormalizeMimeType(codec.MimeType), info.StreamID, ssrc, info.RepairSSRC)
}
}
}
+7 -12
View File
@@ -184,11 +184,6 @@ func (w wrappedICECandidatePairLogger) MarshalLogObject(e zapcore.ObjectEncoder)
// -------------------------------------------------------------------
type SimulcastTrackInfo struct {
Mid string
Rid string
}
type trackDescription struct {
mid string
sender *webrtc.RTPSender
@@ -299,7 +294,7 @@ type TransportParams struct {
EnabledCodecs []*livekit.Codec
Logger logger.Logger
Transport livekit.SignalTarget
SimTracks map[uint32]SimulcastTrackInfo
SimTracks map[uint32]sfuinterceptor.SimulcastTrackInfo
ClientInfo ClientInfo
IsOfferer bool
IsSendSide bool
@@ -455,7 +450,7 @@ func newPeerConnection(
ir.Add(lkinterceptor.NewRTTFromXRFactory(func(rtt uint32) {}))
}
if len(params.SimTracks) > 0 {
f, err := NewUnhandleSimulcastInterceptorFactory(UnhandleSimulcastTracks(params.SimTracks))
f, err := sfuinterceptor.NewUnhandleSimulcastInterceptorFactory(sfuinterceptor.UnhandleSimulcastTracks(params.Logger, params.SimTracks))
if err != nil {
params.Logger.Warnw("NewUnhandleSimulcastInterceptorFactory failed", err)
} else {
@@ -498,9 +493,9 @@ func newPeerConnection(
}
rtxInfoExtractorFactory := sfuinterceptor.NewRTXInfoExtractorFactory(
setTWCCForVideo,
func(repair, base uint32) {
params.Logger.Debugw("rtx pair found from extension", "repair", repair, "base", base)
params.Config.BufferFactory.SetRTXPair(repair, base)
func(repair, base uint32, rsid string) {
params.Logger.Debugw("rtx pair found from extension", "repair", repair, "base", base, "rsid", rsid)
params.Config.BufferFactory.SetRTXPair(repair, base, rsid)
},
params.Logger,
)
@@ -1628,7 +1623,7 @@ func (t *PCTransport) HandleRemoteDescription(sd webrtc.SessionDescription, remo
if len(rtxRepairs) > 0 {
t.params.Logger.Debugw("rtx pairs found from sdp", "ssrcs", rtxRepairs)
for repair, base := range rtxRepairs {
t.params.Config.BufferFactory.SetRTXPair(repair, base)
t.params.Config.BufferFactory.SetRTXPair(repair, base, "")
}
}
return nil
@@ -2837,7 +2832,7 @@ func (t *PCTransport) handleRemoteOfferReceived(sd *webrtc.SessionDescription, o
if len(rtxRepairs) > 0 {
t.params.Logger.Debugw("rtx pairs found from sdp", "ssrcs", rtxRepairs)
for repair, base := range rtxRepairs {
t.params.Config.BufferFactory.SetRTXPair(repair, base)
t.params.Config.BufferFactory.SetRTXPair(repair, base, "")
}
}
+2 -1
View File
@@ -39,6 +39,7 @@ import (
"github.com/livekit/livekit-server/pkg/rtc/types"
"github.com/livekit/livekit-server/pkg/sfu"
"github.com/livekit/livekit-server/pkg/sfu/datachannel"
"github.com/livekit/livekit-server/pkg/sfu/interceptor"
"github.com/livekit/livekit-server/pkg/sfu/pacer"
"github.com/livekit/livekit-server/pkg/telemetry"
)
@@ -80,7 +81,7 @@ type TransportManagerParams struct {
CongestionControlConfig config.CongestionControlConfig
EnabledSubscribeCodecs []*livekit.Codec
EnabledPublishCodecs []*livekit.Codec
SimTracks map[uint32]SimulcastTrackInfo
SimTracks map[uint32]interceptor.SimulcastTrackInfo
ClientInfo ClientInfo
Migration bool
AllowTCPFallback bool
-132
View File
@@ -1,132 +0,0 @@
// Copyright 2023 LiveKit, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package rtc
import (
"github.com/pion/interceptor"
"github.com/pion/rtp"
"github.com/pion/sdp/v3"
"github.com/pion/webrtc/v4"
"github.com/livekit/livekit-server/pkg/sfu/utils"
)
const (
simulcastProbeCount = 10
)
type UnhandleSimulcastOption func(r *UnhandleSimulcastInterceptor) error
func UnhandleSimulcastTracks(tracks map[uint32]SimulcastTrackInfo) UnhandleSimulcastOption {
return func(r *UnhandleSimulcastInterceptor) error {
r.simTracks = tracks
return nil
}
}
type UnhandleSimulcastInterceptorFactory struct {
opts []UnhandleSimulcastOption
}
func (f *UnhandleSimulcastInterceptorFactory) NewInterceptor(id string) (interceptor.Interceptor, error) {
i := &UnhandleSimulcastInterceptor{simTracks: map[uint32]SimulcastTrackInfo{}}
for _, o := range f.opts {
if err := o(i); err != nil {
return nil, err
}
}
return i, nil
}
func NewUnhandleSimulcastInterceptorFactory(opts ...UnhandleSimulcastOption) (*UnhandleSimulcastInterceptorFactory, error) {
return &UnhandleSimulcastInterceptorFactory{opts: opts}, nil
}
type unhandleSimulcastRTPReader struct {
SimulcastTrackInfo
tryTimes int
reader interceptor.RTPReader
midExtensionID uint8
streamIDExtensionID uint8
}
func (r *unhandleSimulcastRTPReader) Read(b []byte, a interceptor.Attributes) (int, interceptor.Attributes, error) {
n, a, err := r.reader.Read(b, a)
if r.tryTimes < 0 || err != nil {
return n, a, err
}
header := rtp.Header{}
hsize, err := header.Unmarshal(b[:n])
if err != nil {
return n, a, nil
}
var mid, rid string
if payload := header.GetExtension(r.midExtensionID); payload != nil {
mid = string(payload)
}
if payload := header.GetExtension(r.streamIDExtensionID); payload != nil {
rid = string(payload)
}
if mid != "" && rid != "" {
r.tryTimes = -1
return n, a, nil
}
r.tryTimes--
if mid == "" {
header.SetExtension(r.midExtensionID, []byte(r.Mid))
}
if rid == "" {
header.SetExtension(r.streamIDExtensionID, []byte(r.Rid))
}
hsize2 := header.MarshalSize()
if hsize2-hsize+n > len(b) { // no enough buf to set extension
return n, a, nil
}
copy(b[hsize2:], b[hsize:n])
header.MarshalTo(b)
return hsize2 - hsize + n, a, nil
}
type UnhandleSimulcastInterceptor struct {
interceptor.NoOp
simTracks map[uint32]SimulcastTrackInfo
}
func (u *UnhandleSimulcastInterceptor) BindRemoteStream(info *interceptor.StreamInfo, reader interceptor.RTPReader) interceptor.RTPReader {
if t, ok := u.simTracks[info.SSRC]; ok {
// if we support fec for simulcast streams at future, should get rsid extensions
midExtensionID := utils.GetHeaderExtensionID(info.RTPHeaderExtensions, webrtc.RTPHeaderExtensionCapability{URI: sdp.SDESMidURI})
streamIDExtensionID := utils.GetHeaderExtensionID(info.RTPHeaderExtensions, webrtc.RTPHeaderExtensionCapability{URI: sdp.SDESRTPStreamIDURI})
if midExtensionID == 0 || streamIDExtensionID == 0 {
return reader
}
return &unhandleSimulcastRTPReader{
SimulcastTrackInfo: t,
reader: reader,
tryTimes: simulcastProbeCount,
midExtensionID: uint8(midExtensionID),
streamIDExtensionID: uint8(streamIDExtensionID),
}
}
return reader
}
+20
View File
@@ -68,6 +68,7 @@ type Buffer struct {
onClose func()
onRtcpFeedback func([]rtcp.Packet)
onFinalRtpStats func(*livekit.RTPStats)
onNotifyRTX func(uint32, uint32, string)
primaryBufferForRTX *Buffer
rtxPktBuf []byte
@@ -219,6 +220,12 @@ func (b *Buffer) SetPrimaryBufferForRTX(primaryBuffer *Buffer) {
}
}
func (b *Buffer) NotifyRTX(ssrc uint32, repairSSRC uint32, rsid string) {
if onNotifyRTX := b.getOnNotifyRTX(); onNotifyRTX != nil {
onNotifyRTX(ssrc, repairSSRC, rsid)
}
}
func (b *Buffer) writeRTX(rtxPkt *rtp.Packet, arrivalTime int64) {
b.Lock()
defer b.Unlock()
@@ -435,3 +442,16 @@ func (b *Buffer) getOnFinalRtpStats() func(*livekit.RTPStats) {
return b.onFinalRtpStats
}
func (b *Buffer) OnNotifyRTX(fn func(ssrc uint32, repairSSRC uint32, rsid string)) {
b.Lock()
b.onNotifyRTX = fn
b.Unlock()
}
func (b *Buffer) getOnNotifyRTX() func(ssrc uint32, repairSSRC uint32, rsid string) {
b.RLock()
defer b.RUnlock()
return b.onNotifyRTX
}
+1 -1
View File
@@ -402,7 +402,7 @@ func (b *BufferBase) CloseWithReason(reason string) (stats *livekit.RTPStats, er
b.StopKeyFrameSeeder()
b.Lock()
b.stopRTPStats(reason)
stats, _ = b.stopRTPStats(reason)
b.readCond.Broadcast()
b.Unlock()
+4 -1
View File
@@ -118,7 +118,7 @@ func (f *Factory) GetRTCPReader(ssrc uint32) *RTCPReader {
return f.rtcpReaders[ssrc]
}
func (f *Factory) SetRTXPair(repair, base uint32) {
func (f *Factory) SetRTXPair(repair, base uint32, rsid string) {
f.Lock()
repairBuffer, baseBuffer := f.rtpBuffers[repair], f.rtpBuffers[base]
if repairBuffer == nil || baseBuffer == nil {
@@ -127,5 +127,8 @@ func (f *Factory) SetRTXPair(repair, base uint32) {
f.Unlock()
if repairBuffer != nil && baseBuffer != nil {
repairBuffer.SetPrimaryBufferForRTX(baseBuffer)
if rsid != "" {
baseBuffer.NotifyRTX(base, repair, rsid)
}
}
}
+10 -3
View File
@@ -39,13 +39,17 @@ type streamInfo struct {
type RTXInfoExtractorFactory struct {
onStreamFound func(*interceptor.StreamInfo)
onRTXPairFound func(repair, base uint32)
onRTXPairFound func(repair, base uint32, rsid string)
lock sync.Mutex
streams map[uint32]streamInfo
logger logger.Logger
}
func NewRTXInfoExtractorFactory(onStreamFound func(*interceptor.StreamInfo), onRTXPairFound func(repair, base uint32), logger logger.Logger) *RTXInfoExtractorFactory {
func NewRTXInfoExtractorFactory(
onStreamFound func(*interceptor.StreamInfo),
onRTXPairFound func(repair, base uint32, rsid string),
logger logger.Logger,
) *RTXInfoExtractorFactory {
return &RTXInfoExtractorFactory{
onStreamFound: onStreamFound,
onRTXPairFound: onRTXPairFound,
@@ -63,6 +67,7 @@ func (f *RTXInfoExtractorFactory) NewInterceptor(id string) (interceptor.Interce
func (f *RTXInfoExtractorFactory) SetStreamInfo(ssrc uint32, mid, rid, rsid string) {
var repairSsrc, baseSsrc uint32
var repairSid string
f.lock.Lock()
if mid == "" || (rid == "" && rsid == "") {
@@ -76,6 +81,7 @@ func (f *RTXInfoExtractorFactory) SetStreamInfo(ssrc uint32, mid, rid, rsid stri
if info.mid == mid && info.rid == rsid {
repairSsrc = ssrc
baseSsrc = base
repairSid = rsid
delete(f.streams, base)
break
}
@@ -86,6 +92,7 @@ func (f *RTXInfoExtractorFactory) SetStreamInfo(ssrc uint32, mid, rid, rsid stri
if info.mid == mid && info.rsid == rid {
repairSsrc = repair
baseSsrc = ssrc
repairSid = info.rid
delete(f.streams, repair)
break
}
@@ -104,7 +111,7 @@ func (f *RTXInfoExtractorFactory) SetStreamInfo(ssrc uint32, mid, rid, rsid stri
f.lock.Unlock()
if repairSsrc != 0 && baseSsrc != 0 {
f.onRTXPairFound(repairSsrc, baseSsrc)
f.onRTXPairFound(repairSsrc, baseSsrc, repairSid)
}
}
+183
View File
@@ -0,0 +1,183 @@
// Copyright 2023 LiveKit, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package interceptor
import (
"github.com/pion/interceptor"
"github.com/pion/rtp"
"github.com/pion/sdp/v3"
"github.com/pion/webrtc/v4"
"go.uber.org/zap/zapcore"
"github.com/livekit/livekit-server/pkg/sfu/utils"
"github.com/livekit/protocol/logger"
)
const (
simulcastProbeCount = 10
)
type SimulcastTrackInfo struct {
Mid string
StreamID string
RepairSSRC uint32 // set only when `IsRepairStream: false`, i. e. RTX SSRC for the primary stream
IsRepairStream bool
}
func (s *SimulcastTrackInfo) MarshalLogObject(e zapcore.ObjectEncoder) error {
e.AddString("Mid", s.Mid)
e.AddString("StreamID", s.StreamID)
e.AddUint32("RepairSSRC", s.RepairSSRC)
e.AddBool("IsRepairStream", s.IsRepairStream)
return nil
}
// -------------------------------------------------------------------
type UnhandleSimulcastOption func(u *UnhandleSimulcastInterceptor) error
func UnhandleSimulcastTracks(logger logger.Logger, tracks map[uint32]SimulcastTrackInfo) UnhandleSimulcastOption {
return func(u *UnhandleSimulcastInterceptor) error {
u.logger = logger
u.simTracks = tracks
return nil
}
}
type UnhandleSimulcastInterceptorFactory struct {
opts []UnhandleSimulcastOption
}
func (f *UnhandleSimulcastInterceptorFactory) NewInterceptor(id string) (interceptor.Interceptor, error) {
i := &UnhandleSimulcastInterceptor{simTracks: map[uint32]SimulcastTrackInfo{}}
for _, o := range f.opts {
if err := o(i); err != nil {
return nil, err
}
}
return i, nil
}
func NewUnhandleSimulcastInterceptorFactory(opts ...UnhandleSimulcastOption) (*UnhandleSimulcastInterceptorFactory, error) {
return &UnhandleSimulcastInterceptorFactory{opts: opts}, nil
}
type unhandleSimulcastRTPReader struct {
SimulcastTrackInfo
logger logger.Logger
tryTimes int
reader interceptor.RTPReader
midExtID uint8
ridExtID uint8
rsidExtID uint8
}
func (u *unhandleSimulcastRTPReader) Read(b []byte, a interceptor.Attributes) (int, interceptor.Attributes, error) {
n, a, err := u.reader.Read(b, a)
if u.tryTimes < 0 || err != nil {
return n, a, err
}
header := rtp.Header{}
hsize, err := header.Unmarshal(b[:n])
if err != nil {
return n, a, nil
}
var mid, rid, rsid string
if payload := header.GetExtension(u.midExtID); payload != nil {
mid = string(payload)
}
if payload := header.GetExtension(u.ridExtID); payload != nil {
rid = string(payload)
}
if payload := header.GetExtension(u.rsidExtID); payload != nil {
rid = string(payload)
}
if mid != "" && (rid != "" || rsid != "") {
u.logger.Debugw(
"unhandle stream found",
"mid", mid,
"rid", rid,
"rsid", rsid,
"ssrc", header.SSRC,
"simulcastTrackInfo", u.SimulcastTrackInfo,
)
u.tryTimes = -1
return n, a, nil
} else {
// ignore padding only packet for probe count
if !(header.Padding && n-header.MarshalSize()-int(b[n-1]) == 0) {
u.tryTimes--
}
}
if mid == "" {
header.SetExtension(u.midExtID, []byte(u.Mid))
}
if rid == "" && !u.IsRepairStream {
header.SetExtension(u.ridExtID, []byte(u.StreamID))
}
if rsid == "" && u.IsRepairStream {
header.SetExtension(u.rsidExtID, []byte(u.StreamID))
}
hsize2 := header.MarshalSize()
if hsize2-hsize+n > len(b) { // no enough buf to set extension
return n, a, nil
}
copy(b[hsize2:], b[hsize:n])
header.MarshalTo(b)
u.logger.Debugw(
"unhandle stream injecting",
"mid", mid,
"rid", rid,
"rsid", rsid,
"ssrc", header.SSRC,
"simulcastTrackInfo", u.SimulcastTrackInfo,
)
return hsize2 - hsize + n, a, nil
}
type UnhandleSimulcastInterceptor struct {
interceptor.NoOp
logger logger.Logger
simTracks map[uint32]SimulcastTrackInfo
}
func (u *UnhandleSimulcastInterceptor) BindRemoteStream(info *interceptor.StreamInfo, reader interceptor.RTPReader) interceptor.RTPReader {
if t, ok := u.simTracks[info.SSRC]; ok {
midExtensionID := utils.GetHeaderExtensionID(info.RTPHeaderExtensions, webrtc.RTPHeaderExtensionCapability{URI: sdp.SDESMidURI})
streamIDExtensionID := utils.GetHeaderExtensionID(info.RTPHeaderExtensions, webrtc.RTPHeaderExtensionCapability{URI: sdp.SDESRTPStreamIDURI})
repairStreamIDExtensionID := utils.GetHeaderExtensionID(info.RTPHeaderExtensions, webrtc.RTPHeaderExtensionCapability{URI: sdp.SDESRepairRTPStreamIDURI})
if midExtensionID == 0 || streamIDExtensionID == 0 || repairStreamIDExtensionID == 0 {
return reader
}
return &unhandleSimulcastRTPReader{
SimulcastTrackInfo: t,
logger: u.logger,
reader: reader,
tryTimes: simulcastProbeCount,
midExtID: uint8(midExtensionID),
ridExtID: uint8(streamIDExtensionID),
rsidExtID: uint8(repairStreamIDExtensionID),
}
}
return reader
}