diff --git a/go.mod b/go.mod index d328cc385..9a749091e 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index c24812ec2..fa294f7ab 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/pkg/rtc/mediatrack.go b/pkg/rtc/mediatrack.go index 5d3a5b1f4..8fd5de1f9 100644 --- a/pkg/rtc/mediatrack.go +++ b/pkg/rtc/mediatrack.go @@ -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() diff --git a/pkg/rtc/mediatrackreceiver.go b/pkg/rtc/mediatrackreceiver.go index b43e3d9c7..8bbbbb2d1 100644 --- a/pkg/rtc/mediatrackreceiver.go +++ b/pkg/rtc/mediatrackreceiver.go @@ -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 } } diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 55b87165b..c8735e44a 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -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) } } } diff --git a/pkg/rtc/transport.go b/pkg/rtc/transport.go index e3dcd4810..e2bcd416c 100644 --- a/pkg/rtc/transport.go +++ b/pkg/rtc/transport.go @@ -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, "") } } diff --git a/pkg/rtc/transportmanager.go b/pkg/rtc/transportmanager.go index 9d0e90ccf..c3c01b7b9 100644 --- a/pkg/rtc/transportmanager.go +++ b/pkg/rtc/transportmanager.go @@ -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 diff --git a/pkg/rtc/unhandlesimulcast.go b/pkg/rtc/unhandlesimulcast.go deleted file mode 100644 index 3146cba32..000000000 --- a/pkg/rtc/unhandlesimulcast.go +++ /dev/null @@ -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 -} diff --git a/pkg/sfu/buffer/buffer.go b/pkg/sfu/buffer/buffer.go index dc45bb0d6..e904cb330 100644 --- a/pkg/sfu/buffer/buffer.go +++ b/pkg/sfu/buffer/buffer.go @@ -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 +} diff --git a/pkg/sfu/buffer/buffer_base.go b/pkg/sfu/buffer/buffer_base.go index f4d525111..d2056f43b 100644 --- a/pkg/sfu/buffer/buffer_base.go +++ b/pkg/sfu/buffer/buffer_base.go @@ -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() diff --git a/pkg/sfu/buffer/factory.go b/pkg/sfu/buffer/factory.go index ae0bb0733..c792eeeca 100644 --- a/pkg/sfu/buffer/factory.go +++ b/pkg/sfu/buffer/factory.go @@ -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) + } } } diff --git a/pkg/sfu/interceptor/rtx.go b/pkg/sfu/interceptor/rtx.go index 7937c1545..47bf5989d 100644 --- a/pkg/sfu/interceptor/rtx.go +++ b/pkg/sfu/interceptor/rtx.go @@ -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) } } diff --git a/pkg/sfu/interceptor/unhandlesimulcast.go b/pkg/sfu/interceptor/unhandlesimulcast.go new file mode 100644 index 000000000..33fbd133a --- /dev/null +++ b/pkg/sfu/interceptor/unhandlesimulcast.go @@ -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 +}