From 7c8ea115053e15ae80f57f9997c25b2a47603b92 Mon Sep 17 00:00:00 2001 From: Raja Subramanian Date: Tue, 23 Dec 2025 21:35:48 +0530 Subject: [PATCH] Refactor receiver and buffer into Base and higher layer. (#4185) * Refactor receiver and buffer into Base and higher layer. To be able to share code/functionality with relay. * WIP * WIP * WIP * WIP * WIP * WIP * WIP * WIP * clean up * deps * fix test * fix test --- go.mod | 14 +- go.sum | 30 +- pkg/rtc/mediatrack.go | 13 +- pkg/rtc/wrappedreceiver.go | 2 +- pkg/sfu/buffer/buffer.go | 1270 ++------------------- pkg/sfu/buffer/buffer_base.go | 1493 +++++++++++++++++++++++++ pkg/sfu/receiver.go | 847 ++------------ pkg/sfu/receiver_base.go | 1111 ++++++++++++++++++ pkg/sfu/redreceiver_test.go | 88 +- pkg/sfu/rtpstats/rtpstats_base.go | 2 +- pkg/sfu/rtpstats/rtpstats_receiver.go | 4 + pkg/sfu/rtpstats/rtpstats_sender.go | 4 +- pkg/sfu/track_remote.go | 14 + 13 files changed, 2877 insertions(+), 2015 deletions(-) create mode 100644 pkg/sfu/buffer/buffer_base.go create mode 100644 pkg/sfu/receiver_base.go diff --git a/go.mod b/go.mod index 22f00d02c..63de43d85 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.20251222031942-17875e94bb49 + github.com/livekit/protocol v1.43.5-0.20251222225221-fa169ac100d9 github.com/livekit/psrpc v0.7.1 github.com/mackerelio/go-osstat v0.2.6 github.com/magefile/mage v1.15.0 @@ -36,9 +36,9 @@ require ( github.com/pion/ice/v4 v4.1.0 github.com/pion/interceptor v0.1.42 github.com/pion/rtcp v1.2.16 - github.com/pion/rtp v1.8.26 + github.com/pion/rtp v1.8.27 github.com/pion/sctp v1.8.41 - github.com/pion/sdp/v3 v3.0.16 + github.com/pion/sdp/v3 v3.0.17 github.com/pion/transport/v3 v3.1.1 github.com/pion/turn/v4 v4.1.3 github.com/pion/webrtc/v4 v4.1.8 @@ -55,7 +55,7 @@ require ( go.uber.org/atomic v1.11.0 go.uber.org/multierr v1.11.0 go.uber.org/zap v1.27.1 - golang.org/x/exp v0.0.0-20251209150349-8475f28825e9 + golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 golang.org/x/mod v0.31.0 golang.org/x/sync v0.19.0 google.golang.org/protobuf v1.36.11 @@ -145,8 +145,8 @@ require ( golang.org/x/sys v0.39.0 // indirect golang.org/x/text v0.32.0 // indirect golang.org/x/tools v0.40.0 // indirect - google.golang.org/genproto/googleapis/api v0.0.0-20251213004720-97cd9d5aeac2 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20251213004720-97cd9d5aeac2 // indirect - google.golang.org/grpc v1.77.0 // indirect + google.golang.org/genproto/googleapis/api v0.0.0-20251222181119-0a764e51fe1b // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b // indirect + google.golang.org/grpc v1.78.0 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect ) diff --git a/go.sum b/go.sum index bd1014aac..cdeb71fb1 100644 --- a/go.sum +++ b/go.sum @@ -173,10 +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.20251217174542-5c369b107325 h1:oaqg6YhJDFh1JShsEUK7iXlkLFcm6GEiaVhm0JP+Rsg= -github.com/livekit/protocol v1.43.5-0.20251217174542-5c369b107325/go.mod h1:n00Ul4P6o2YILGhxw+O57B0h/bF3Je9PzRN36fElCmw= -github.com/livekit/protocol v1.43.5-0.20251222031942-17875e94bb49 h1:c2mvLX0IQE7cC/LYYtBeKUrGdHVuXAJMtj4fcjN9vTk= -github.com/livekit/protocol v1.43.5-0.20251222031942-17875e94bb49/go.mod h1:n00Ul4P6o2YILGhxw+O57B0h/bF3Je9PzRN36fElCmw= +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/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= @@ -256,12 +254,12 @@ github.com/pion/randutil v0.1.0 h1:CFG1UdESneORglEsnimhUjf33Rwjubwj6xfiOXBa3mA= github.com/pion/randutil v0.1.0/go.mod h1:XcJrSMMbbMRhASFVOlj/5hQial/Y8oH/HVo7TBZq+j8= github.com/pion/rtcp v1.2.16 h1:fk1B1dNW4hsI78XUCljZJlC4kZOPk67mNRuQ0fcEkSo= github.com/pion/rtcp v1.2.16/go.mod h1:/as7VKfYbs5NIb4h6muQ35kQF/J0ZVNz2Z3xKoCBYOo= -github.com/pion/rtp v1.8.26 h1:VB+ESQFQhBXFytD+Gk8cxB6dXeVf2WQzg4aORvAvAAc= -github.com/pion/rtp v1.8.26/go.mod h1:rF5nS1GqbR7H/TCpKwylzeq6yDM+MM6k+On5EgeThEM= +github.com/pion/rtp v1.8.27 h1:kbWTdZr62RDlYjatVAW4qFwrAu9XcGnwMsofCfAHlOU= +github.com/pion/rtp v1.8.27/go.mod h1:rF5nS1GqbR7H/TCpKwylzeq6yDM+MM6k+On5EgeThEM= github.com/pion/sctp v1.8.41 h1:20R4OHAno4Vky3/iE4xccInAScAa83X6nWUfyc65MIs= github.com/pion/sctp v1.8.41/go.mod h1:2wO6HBycUH7iCssuGyc2e9+0giXVW0pyCv3ZuL8LiyY= -github.com/pion/sdp/v3 v3.0.16 h1:0dKzYO6gTAvuLaAKQkC02eCPjMIi4NuAr/ibAwrGDCo= -github.com/pion/sdp/v3 v3.0.16/go.mod h1:9tyKzznud3qiweZcD86kS0ff1pGYB3VX+Bcsmkx6IXo= +github.com/pion/sdp/v3 v3.0.17 h1:9SfLAW/fF1XC8yRqQ3iWGzxkySxup4k4V7yN8Fs8nuo= +github.com/pion/sdp/v3 v3.0.17/go.mod h1:9tyKzznud3qiweZcD86kS0ff1pGYB3VX+Bcsmkx6IXo= github.com/pion/srtp/v3 v3.0.9 h1:lRGF4G61xxj+m/YluB3ZnBpiALSri2lTzba0kGZMrQY= github.com/pion/srtp/v3 v3.0.9/go.mod h1:E+AuWd7Ug2Fp5u38MKnhduvpVkveXJX6J4Lq4rxUYt8= github.com/pion/stun/v3 v3.0.2 h1:BJuGEN2oLrJisiNEJtUTJC4BGbzbfp37LizfqswblFU= @@ -375,8 +373,8 @@ golang.org/x/crypto v0.18.0/go.mod h1:R0j02AL6hcrfOiy9T4ZYp/rcWeMxM3L6QYxlOuEG1m golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU= golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU= golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0= -golang.org/x/exp v0.0.0-20251209150349-8475f28825e9 h1:MDfG8Cvcqlt9XXrmEiD4epKn7VJHZO84hejP9Jmp0MM= -golang.org/x/exp v0.0.0-20251209150349-8475f28825e9/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU= +golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 h1:fQsdNF2N+/YewlRZiricy4P1iimyPKZ/xwniHj8Q2a0= +golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU= golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= @@ -487,12 +485,12 @@ golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8T golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= -google.golang.org/genproto/googleapis/api v0.0.0-20251213004720-97cd9d5aeac2 h1:7LRqPCEdE4TP4/9psdaB7F2nhZFfBiGJomA5sojLWdU= -google.golang.org/genproto/googleapis/api v0.0.0-20251213004720-97cd9d5aeac2/go.mod h1:+rXWjjaukWZun3mLfjmVnQi18E1AsFbDN9QdJ5YXLto= -google.golang.org/genproto/googleapis/rpc v0.0.0-20251213004720-97cd9d5aeac2 h1:2I6GHUeJ/4shcDpoUlLs/2WPnhg7yJwvXtqcMJt9liA= -google.golang.org/genproto/googleapis/rpc v0.0.0-20251213004720-97cd9d5aeac2/go.mod h1:7i2o+ce6H/6BluujYR+kqX3GKH+dChPTQU19wjRPiGk= -google.golang.org/grpc v1.77.0 h1:wVVY6/8cGA6vvffn+wWK5ToddbgdU3d8MNENr4evgXM= -google.golang.org/grpc v1.77.0/go.mod h1:z0BY1iVj0q8E1uSQCjL9cppRj+gnZjzDnzV0dHhrNig= +google.golang.org/genproto/googleapis/api v0.0.0-20251222181119-0a764e51fe1b h1:uA40e2M6fYRBf0+8uN5mLlqUtV192iiksiICIBkYJ1E= +google.golang.org/genproto/googleapis/api v0.0.0-20251222181119-0a764e51fe1b/go.mod h1:Xa7le7qx2vmqB/SzWUBa7KdMjpdpAHlh5QCSnjessQk= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b h1:Mv8VFug0MP9e5vUxfBcE3vUkV6CImK3cMNMIDFjmzxU= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b/go.mod h1:j9x/tPzZkyxcgEFkiKEEGxfvyumM01BEtsW8xzOahRQ= +google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc= +google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= diff --git a/pkg/rtc/mediatrack.go b/pkg/rtc/mediatrack.go index ae95a2601..5d3a5b1f4 100644 --- a/pkg/rtc/mediatrack.go +++ b/pkg/rtc/mediatrack.go @@ -306,7 +306,12 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track sfu.TrackRe case *rtcp.SourceDescription: case *rtcp.SenderReport: if pkt.SSRC == uint32(track.SSRC()) { - buff.SetSenderReportData(pkt.RTPTime, pkt.NTPTime, pkt.PacketCount, pkt.OctetCount) + buff.SetSenderReportData(&livekit.RTCPSenderReportState{ + RtpTimestamp: pkt.RTPTime, + NtpTimestamp: pkt.NTPTime, + Packets: pkt.PacketCount, + Octets: uint64(pkt.OctetCount), + }) } case *rtcp.ExtendedReport: rttFromXR: @@ -527,12 +532,12 @@ func (t *MediaTrack) AddReceiver(receiver *webrtc.RTPReceiver, track sfu.TrackRe return newCodec, false } - var bitrates int + var expectedBitrate int layers := buffer.GetVideoLayersForMimeType(mimeType, ti) if layer >= 0 && len(layers) > int(layer) { - bitrates = int(layers[layer].GetBitrate()) + expectedBitrate = int(layers[layer].GetBitrate()) } - if err := buff.Bind(receiver.GetParameters(), track.Codec().RTPCodecCapability, bitrates); err != nil { + if err := buff.Bind(receiver.GetParameters(), track.Codec().RTPCodecCapability, expectedBitrate); err != nil { t.params.Logger.Warnw( "binding buffer failed", err, "rid", track.RID(), diff --git a/pkg/rtc/wrappedreceiver.go b/pkg/rtc/wrappedreceiver.go index a26205ca3..12784771a 100644 --- a/pkg/rtc/wrappedreceiver.go +++ b/pkg/rtc/wrappedreceiver.go @@ -396,7 +396,7 @@ func (d *DummyReceiver) GetDownTracks() []sfu.TrackSender { return maps.Values(d.downTracks) } -func (d *DummyReceiver) DebugInfo() map[string]interface{} { +func (d *DummyReceiver) DebugInfo() map[string]any { if receiver := d.getReceiver(); receiver != nil { return receiver.DebugInfo() } diff --git a/pkg/sfu/buffer/buffer.go b/pkg/sfu/buffer/buffer.go index abbce1079..5f62ef42a 100644 --- a/pkg/sfu/buffer/buffer.go +++ b/pkg/sfu/buffer/buffer.go @@ -17,47 +17,21 @@ package buffer import ( "encoding/binary" "errors" - "fmt" "io" - "strings" - "sync" - "time" - "github.com/gammazero/deque" "github.com/pion/rtcp" "github.com/pion/rtp" - "github.com/pion/rtp/codecs" - "github.com/pion/sdp/v3" "github.com/pion/webrtc/v4" - "go.uber.org/atomic" - "github.com/livekit/livekit-server/pkg/sfu/audio" - "github.com/livekit/livekit-server/pkg/sfu/mime" - act "github.com/livekit/livekit-server/pkg/sfu/rtpextension/abscapturetime" - dd "github.com/livekit/livekit-server/pkg/sfu/rtpextension/dependencydescriptor" - "github.com/livekit/livekit-server/pkg/sfu/rtpstats" - "github.com/livekit/livekit-server/pkg/sfu/utils" sutils "github.com/livekit/livekit-server/pkg/utils" "github.com/livekit/mediatransportutil/pkg/bucket" - "github.com/livekit/mediatransportutil/pkg/nack" "github.com/livekit/mediatransportutil/pkg/twcc" "github.com/livekit/protocol/livekit" - "github.com/livekit/protocol/logger" "github.com/livekit/protocol/utils/mono" ) -var ( - ExtPacketFactory = &sync.Pool{ - New: func() any { - return &ExtPacket{} - }, - } -) - -// -------------------------------------- - const ( - ReportDelta = 1e9 + rtcpReceiverReportDelta = 1e9 InitPacketBufferSizeVideo = 300 InitPacketBufferSizeAudio = 70 @@ -67,148 +41,51 @@ var ( errInvalidCodec = errors.New("invalid codec") ) +var _ BufferProvider = (*Buffer)(nil) + type pendingPacket struct { arrivalTime int64 packet []byte } -type ExtPacket struct { - VideoLayer - Arrival int64 - ExtSequenceNumber uint64 - ExtTimestamp uint64 - Packet *rtp.Packet - Payload any - KeyFrame bool - RawPacket []byte - DependencyDescriptor *ExtDependencyDescriptor - AbsCaptureTimeExt *act.AbsCaptureTime - IsOutOfOrder bool - IsBuffered bool - IsRestart bool -} - -// VideoSize represents video resolution -type VideoSize struct { - Width uint32 - Height uint32 -} - // Buffer contains all packets type Buffer struct { - sync.RWMutex - readCond *sync.Cond - bucket *bucket.Bucket[uint64, uint16] - nacker *nack.NackQueue - maxVideoPkts int - maxAudioPkts int - codecType webrtc.RTPCodecType - extPackets deque.Deque[*ExtPacket] - pPackets []pendingPacket - closeOnce sync.Once - mediaSSRC uint32 - clockRate uint32 - lastReport int64 - twccExtID uint8 - audioLevelExtID uint8 - bound bool - closed atomic.Bool + *BufferBase - rtpParameters webrtc.RTPParameters - payloadType uint8 - rtxPayloadType uint8 - mime mime.MimeType + pPackets []pendingPacket + lastReportAt int64 + isBound bool - snRangeMap *utils.RangeMap[uint64, uint64] + twcc *twcc.Responder + twccExtID uint8 - latestTSForAudioLevelInitialized bool - latestTSForAudioLevel uint32 - - twcc *twcc.Responder - audioLevelParams audio.AudioLevelParams - audioLevel *audio.AudioLevel - enableAudioLossProxying bool - enableStreamRestartDetection bool + enableAudioLossProxying bool + lastFractionLostToReport uint8 // Last fraction lost from subscribers, should report to publisher; Audio only lastPacketRead int - pliThrottle int64 - - rtpStats *rtpstats.RTPStatsReceiver - rrSnapshotId uint32 - deltaStatsSnapshotId uint32 - ppsSnapshotId uint32 - - lastFractionLostToReport uint8 // Last fraction lost from subscribers, should report to publisher; Audio only - // callbacks - onClose func() - onRtcpFeedback func([]rtcp.Packet) - onRtcpSenderReport func() - onFpsChanged func() - onFinalRtpStats func(*livekit.RTPStats) - onCodecChange func(webrtc.RTPCodecParameters) - onVideoSizeChanged func([]VideoSize) - - // video size tracking for multiple spatial layers - currentVideoSize [DefaultMaxLayerSpatial + 1]VideoSize - - // logger - logger logger.Logger - - // dependency descriptor - ddExtID uint8 - ddParser *DependencyDescriptorParser - - paused bool - frameRateCalculator [DefaultMaxLayerSpatial + 1]FrameRateCalculator - frameRateCalculated bool - - packetNotFoundCount atomic.Uint32 - packetTooOldCount atomic.Uint32 - extPacketTooMuchCount atomic.Uint32 + onClose func() + onRtcpFeedback func([]rtcp.Packet) + onFinalRtpStats func(*livekit.RTPStats) primaryBufferForRTX *Buffer rtxPktBuf []byte - - absCaptureTimeExtID uint8 - - keyFrameSeederGeneration atomic.Int32 } -// NewBuffer constructs a new Buffer func NewBuffer(ssrc uint32, maxVideoPkts, maxAudioPkts int) *Buffer { - l := logger.GetLogger() // will be reset with correct context via SetLogger - b := &Buffer{ - mediaSSRC: ssrc, - maxVideoPkts: maxVideoPkts, - maxAudioPkts: maxAudioPkts, - snRangeMap: utils.NewRangeMap[uint64, uint64](100), - pliThrottle: int64(500 * time.Millisecond), - logger: l.WithComponent(sutils.ComponentPub).WithComponent(sutils.ComponentSFU), - } - b.readCond = sync.NewCond(&b.RWMutex) - b.extPackets.SetBaseCap(128) + b := &Buffer{} + b.BufferBase = NewBufferBase(BufferBaseParams{ + SSRC: ssrc, + MaxVideoPkts: maxVideoPkts, + MaxAudioPkts: maxAudioPkts, + LoggerComponents: []string{sutils.ComponentPub, sutils.ComponentSFU}, + SendPLI: b.sendPLI, + IsReportingEnabled: true, + }) return b } -func (b *Buffer) SetLogger(logger logger.Logger) { - b.Lock() - defer b.Unlock() - - b.logger = logger.WithComponent(sutils.ComponentPub).WithComponent(sutils.ComponentSFU).WithValues("ssrc", b.mediaSSRC) - if b.rtpStats != nil { - b.rtpStats.SetLogger(b.logger) - } -} - -func (b *Buffer) SetPaused(paused bool) { - b.Lock() - defer b.Unlock() - - b.paused = paused -} - func (b *Buffer) SetTWCCAndExtID(twcc *twcc.Responder, extID uint8) { b.Lock() defer b.Unlock() @@ -217,13 +94,6 @@ func (b *Buffer) SetTWCCAndExtID(twcc *twcc.Responder, extID uint8) { b.twccExtID = extID } -func (b *Buffer) SetAudioLevelParams(audioLevelParams audio.AudioLevelParams) { - b.Lock() - defer b.Unlock() - - b.audioLevelParams = audioLevelParams -} - func (b *Buffer) SetAudioLossProxying(enable bool) { b.Lock() defer b.Unlock() @@ -231,177 +101,32 @@ func (b *Buffer) SetAudioLossProxying(enable bool) { b.enableAudioLossProxying = enable } -func (b *Buffer) SetStreamRestartDetection(enable bool) { - b.Lock() - defer b.Unlock() - - b.enableStreamRestartDetection = enable -} - func (b *Buffer) Bind(params webrtc.RTPParameters, codec webrtc.RTPCodecCapability, bitrates int) error { b.Lock() defer b.Unlock() - if b.bound { + if b.isBound { return nil } - b.logger.Debugw("binding track") - if codec.ClockRate == 0 { - b.logger.Warnw("invalid codec", nil, "params", params, "codec", codec, "bitrates", bitrates) - return errInvalidCodec + if err := b.BufferBase.BindLocked(params, codec, bitrates); err != nil { + return err } - b.setupRTPStats(codec.ClockRate) - - b.clockRate = codec.ClockRate - b.lastReport = mono.UnixNano() - b.mime = mime.NormalizeMimeType(codec.MimeType) - b.rtpParameters = params - for _, codecParameter := range params.Codecs { - if mime.IsMimeTypeStringEqual(codecParameter.MimeType, codec.MimeType) { - b.payloadType = uint8(codecParameter.PayloadType) - break - } - } - - if b.payloadType == 0 && !mime.IsMimeTypeStringEqual(codec.MimeType, webrtc.MimeTypePCMU) { - b.logger.Warnw("could not find payload type for codec", nil, "codec", codec.MimeType, "parameters", params) - b.payloadType = uint8(params.Codecs[0].PayloadType) - } - - // find RTX payload type - for _, codec := range params.Codecs { - if mime.IsMimeTypeStringRTX(codec.MimeType) && strings.Contains(codec.SDPFmtpLine, fmt.Sprintf("apt=%d", b.payloadType)) { - b.rtxPayloadType = uint8(codec.PayloadType) - break - } - } - - for _, ext := range params.HeaderExtensions { - switch ext.URI { - case dd.ExtensionURI: - if b.ddExtID != 0 { - b.logger.Warnw("multiple dependency descriptor extensions found", nil, "id", ext.ID, "previous", b.ddExtID) - continue - } - b.ddExtID = uint8(ext.ID) - b.createDDParserAndFrameRateCalculator() - - case sdp.AudioLevelURI: - b.audioLevelExtID = uint8(ext.ID) - b.audioLevel = audio.NewAudioLevel(b.audioLevelParams) - - case act.AbsCaptureTimeURI: - b.absCaptureTimeExtID = uint8(ext.ID) - } - } - - switch { - case mime.IsMimeTypeAudio(b.mime): - b.codecType = webrtc.RTPCodecTypeAudio - b.bucket = bucket.NewBucket[uint64, uint16](InitPacketBufferSizeAudio, bucket.RTPMaxPktSize, bucket.RTPSeqNumOffset) - - case mime.IsMimeTypeVideo(b.mime): - b.codecType = webrtc.RTPCodecTypeVideo - b.bucket = bucket.NewBucket[uint64, uint16](InitPacketBufferSizeVideo, bucket.RTPMaxPktSize, bucket.RTPSeqNumOffset) - if b.frameRateCalculator[0] == nil { - b.createFrameRateCalculator() - } - if bitrates > 0 { - pps := bitrates / 8 / 1200 - for pps > b.bucket.Capacity() { - if b.bucket.Grow() >= b.maxVideoPkts { - break - } - } - } - - default: - b.codecType = webrtc.RTPCodecType(0) - } - - for _, fb := range codec.RTCPFeedback { - switch fb.Type { - case webrtc.TypeRTCPFBGoogREMB: - b.logger.Debugw("Setting feedback", "type", webrtc.TypeRTCPFBGoogREMB) - b.logger.Debugw("REMB not supported, RTCP feedback will not be generated") - case webrtc.TypeRTCPFBNACK: - // pion use a single mediaengine to manage negotiated codecs of peerconnection, that means we can't have different - // codec settings at track level for same codec type, so enable nack for all audio receivers but don't create nack queue - // for red codec. - if b.mime == mime.MimeTypeRED { - break - } - b.logger.Debugw("Setting feedback", "type", webrtc.TypeRTCPFBNACK) - b.nacker = nack.NewNACKQueue(nack.NackQueueParamsDefault) - } - } + b.lastReportAt = mono.UnixNano() if len(b.pPackets) != 0 { b.logger.Debugw("releasing queued packets on bind", "count", len(b.pPackets)) } for _, pp := range b.pPackets { - b.calc(pp.packet, nil, pp.arrivalTime, false, true) + b.calc(pp.packet, nil, pp.arrivalTime, true, false) } b.pPackets = nil - b.bound = true - if mime.IsMimeTypeVideo(b.mime) { - go b.seedKeyFrame(b.keyFrameSeederGeneration.Inc()) - } + b.isBound = true return nil } -func (b *Buffer) OnCodecChange(fn func(webrtc.RTPCodecParameters)) { - b.Lock() - b.onCodecChange = fn - b.Unlock() -} - -func (b *Buffer) setupRTPStats(clockRate uint32) { - b.rtpStats = rtpstats.NewRTPStatsReceiver(rtpstats.RTPStatsParams{ - ClockRate: clockRate, - Logger: b.logger, - }) - b.rrSnapshotId = b.rtpStats.NewSnapshotId() - b.deltaStatsSnapshotId = b.rtpStats.NewSnapshotId() - b.ppsSnapshotId = b.rtpStats.NewSnapshotId() -} - -func (b *Buffer) createDDParserAndFrameRateCalculator() { - if mime.IsMimeTypeSVCCapable(b.mime) || b.mime == mime.MimeTypeVP8 { - frc := NewFrameRateCalculatorDD(b.clockRate, b.logger) - for i := range b.frameRateCalculator { - b.frameRateCalculator[i] = frc.GetFrameRateCalculatorForSpatial(int32(i)) - } - b.ddParser = NewDependencyDescriptorParser( - b.ddExtID, - b.logger, - func(spatial, temporal int32) { - frc.SetMaxLayer(spatial, temporal) - }, - false, - ) - } -} - -func (b *Buffer) createFrameRateCalculator() { - switch b.mime { - case mime.MimeTypeVP8: - b.frameRateCalculator[0] = NewFrameRateCalculatorVP8(b.clockRate, b.logger) - - case mime.MimeTypeVP9: - frc := NewFrameRateCalculatorVP9(b.clockRate, b.logger) - for i := range b.frameRateCalculator { - b.frameRateCalculator[i] = frc.GetFrameRateCalculatorForSpatial(int32(i)) - } - - case mime.MimeTypeH265: - b.frameRateCalculator[0] = NewFrameRateCalculatorH26x(b.clockRate, b.logger) - } -} - // Write adds an RTP Packet, ordering is not guaranteed, newer packets may arrive later // //go:noinline @@ -413,14 +138,14 @@ func (b *Buffer) Write(pkt []byte) (n int, err error) { } b.Lock() - if b.closed.Load() { + if b.BufferBase.IsClosed() { b.Unlock() err = io.EOF return } now := mono.UnixNano() - if b.twcc != nil && b.twccExtID != 0 && !b.closed.Load() { + if b.twcc != nil && b.twccExtID != 0 { if ext := rtpPacket.GetExtension(b.twccExtID); ext != nil { b.twcc.Push(rtpPacket.SSRC, binary.BigEndian.Uint16(ext[0:2]), now, rtpPacket.Marker) } @@ -446,7 +171,7 @@ func (b *Buffer) Write(pkt []byte) (n int, err error) { return } - if !b.bound { + if !b.isBound { packet := make([]byte, len(pkt)) copy(packet, pkt) @@ -455,7 +180,7 @@ func (b *Buffer) Write(pkt []byte) (n int, err error) { } startIdx := 0 - overflow := len(b.pPackets) - max(b.maxVideoPkts, b.maxAudioPkts) + overflow := len(b.pPackets) - max(b.BufferBase.MaxVideoPkts(), b.BufferBase.MaxAudioPkts()) if overflow > 0 { startIdx = overflow } @@ -464,13 +189,12 @@ func (b *Buffer) Write(pkt []byte) (n int, err error) { arrivalTime: now, }) - b.readCond.Broadcast() + b.BufferBase.NotifyRead() b.Unlock() return } b.calc(pkt, &rtpPacket, now, false, false) - b.readCond.Broadcast() b.Unlock() return } @@ -481,6 +205,7 @@ func (b *Buffer) SetPrimaryBufferForRTX(primaryBuffer *Buffer) { pkts := b.pPackets b.pPackets = nil b.Unlock() + for _, pp := range pkts { var rtpPacket rtp.Packet err := rtpPacket.Unmarshal(pp.packet) @@ -497,7 +222,7 @@ func (b *Buffer) SetPrimaryBufferForRTX(primaryBuffer *Buffer) { func (b *Buffer) writeRTX(rtxPkt *rtp.Packet, arrivalTime int64) { b.Lock() defer b.Unlock() - if !b.bound { + if !b.isBound { return } @@ -518,25 +243,25 @@ func (b *Buffer) writeRTX(rtxPkt *rtp.Packet, arrivalTime int64) { repairedPkt := *rtxPkt repairedPkt.PayloadType = b.payloadType repairedPkt.SequenceNumber = binary.BigEndian.Uint16(rtxPkt.Payload[:2]) - repairedPkt.SSRC = b.mediaSSRC + repairedPkt.SSRC = b.BufferBase.SSRC() repairedPkt.Payload = rtxPkt.Payload[2:] n, err := repairedPkt.MarshalTo(b.rtxPktBuf) if err != nil { - b.logger.Errorw("could not marshal repaired packet", err, "ssrc", b.mediaSSRC, "sn", repairedPkt.SequenceNumber) + b.logger.Errorw("could not marshal repaired packet", err, "ssrc", b.BufferBase.SSRC(), "sn", repairedPkt.SequenceNumber) return } - b.calc(b.rtxPktBuf[:n], &repairedPkt, arrivalTime, true, false) - b.readCond.Broadcast() + b.calc(b.rtxPktBuf[:n], &repairedPkt, arrivalTime, false, true) } func (b *Buffer) Read(buff []byte) (n int, err error) { b.Lock() for { - if b.closed.Load() { + if b.BufferBase.IsClosed() { b.Unlock() return 0, io.EOF } + if b.pPackets != nil && len(b.pPackets) > b.lastPacketRead { if len(buff) < len(b.pPackets[b.lastPacketRead].packet) { b.Unlock() @@ -548,58 +273,26 @@ func (b *Buffer) Read(buff []byte) (n int, err error) { b.Unlock() return } - b.readCond.Wait() - } -} - -func (b *Buffer) ReadExtended(buf []byte) (*ExtPacket, error) { - b.Lock() - for { - if b.closed.Load() { - b.Unlock() - return nil, io.EOF - } - if b.extPackets.Len() > 0 { - ep := b.extPackets.PopFront() - patched := b.patchExtPacket(ep, buf) - if patched == nil { - ReleaseExtPacket(ep) - continue - } - - b.Unlock() - return patched, nil - } - b.readCond.Wait() + b.BufferBase.WaitRead() } } func (b *Buffer) Close() error { - b.closeOnce.Do(func() { - b.closed.Store(true) + stats, err := b.BufferBase.CloseWithReason("close") + if err != nil { + return err + } - b.RLock() - rtpStats := b.rtpStats - b.readCond.Broadcast() - b.RUnlock() - - if rtpStats != nil { - rtpStats.Stop() - b.logger.Debugw("rtp stats", - "direction", "upstream", - "stats", rtpStats, - ) - if cb := b.getOnFinalRtpStats(); cb != nil { - cb(rtpStats.ToProto()) - } + if stats != nil { + if cb := b.getOnFinalRtpStats(); cb != nil { + cb(stats) } + } - if cb := b.getOnClose(); cb != nil { - cb() - } + if cb := b.getOnClose(); cb != nil { + cb() + } - go b.flushExtPackets() - }) return nil } @@ -616,26 +309,13 @@ func (b *Buffer) getOnClose() func() { return b.onClose } -func (b *Buffer) SetPLIThrottle(duration int64) { - b.Lock() - defer b.Unlock() - - b.pliThrottle = duration -} - -func (b *Buffer) SendPLI(force bool) { - b.RLock() - rtpStats := b.rtpStats - pliThrottle := b.pliThrottle - b.RUnlock() - - if (rtpStats == nil && !force) || !rtpStats.CheckAndUpdatePli(pliThrottle, force) { - return - } - - b.logger.Debugw("send pli", "ssrc", b.mediaSSRC, "force", force) +func (b *Buffer) sendPLI() { + ssrc := b.BufferBase.SSRC() pli := []rtcp.Packet{ - &rtcp.PictureLossIndication{SenderSSRC: b.mediaSSRC, MediaSSRC: b.mediaSSRC}, + &rtcp.PictureLossIndication{ + SenderSSRC: ssrc, + MediaSSRC: ssrc, + }, } if cb := b.getOnRtcpFeedback(); cb != nil { @@ -643,513 +323,48 @@ func (b *Buffer) SendPLI(force bool) { } } -func (b *Buffer) SetRTT(rtt uint32) { - b.Lock() - defer b.Unlock() - - if rtt == 0 { - return - } - - if b.nacker != nil { - b.nacker.SetRTT(rtt) - } - - if b.rtpStats != nil { - b.rtpStats.UpdateRtt(rtt) - } -} - -func (b *Buffer) calc(rawPkt []byte, rtpPacket *rtp.Packet, arrivalTime int64, isRTX bool, isBuffered bool) { - defer func() { - b.doNACKs() - - b.doReports(arrivalTime) - }() - - if rtpPacket == nil { - rtpPacket = &rtp.Packet{} - if err := rtpPacket.Unmarshal(rawPkt); err != nil { - b.logger.Errorw("could not unmarshal RTP packet", err) - return - } - } - - // process header extensions always as padding packets could be used for probing - b.processHeaderExtensions(rtpPacket, arrivalTime, isRTX) - - isRestart := false - flowState := b.updateStreamState(rtpPacket, arrivalTime) - switch flowState.UnhandledReason { - case rtpstats.RTPFlowUnhandledReasonNone: - case rtpstats.RTPFlowUnhandledReasonRestart: - if !b.enableStreamRestartDetection { - return - } - - b.rtpStats.Stop() - b.logger.Infow("stream restart - rtp stats", b.rtpStats) - - b.snRangeMap = utils.NewRangeMap[uint64, uint64](100) - b.setupRTPStats(b.clockRate) - b.bucket.ResyncOnNextPacket() - if b.nacker != nil { - b.nacker = nack.NewNACKQueue(nack.NackQueueParamsDefault) - } - b.flushExtPacketsLocked() - - flowState = b.updateStreamState(rtpPacket, arrivalTime) - isRestart = true - default: - return - } - - if len(rtpPacket.Payload) == 0 && (!flowState.IsOutOfOrder || flowState.IsDuplicate) { - // drop padding only in-order or duplicate packet - if !flowState.IsOutOfOrder { - // in-order packet - increment sequence number offset for subsequent packets - // Example: - // 40 - regular packet - pass through as sequence number 40 - // 41 - missing packet - don't know what it is, could be padding or not - // 42 - padding only packet - in-order - drop - increment sequence number offset to 1 - - // range[0, 42] = 0 offset - // 41 - arrives out of order - get offset 0 from cache - passed through as sequence number 41 - // 43 - regular packet - offset = 1 (running offset) - passes through as sequence number 42 - // 44 - padding only - in order - drop - increment sequence number offset to 2 - // range[0, 42] = 0 offset, range[43, 44] = 1 offset - // 43 - regular packet - out of order + duplicate - offset = 1 from cache - - // adjusted sequence number is 42, will be dropped by RTX buffer AddPacket method as duplicate - // 45 - regular packet - offset = 2 (running offset) - passed through with adjusted sequence number as 43 - // 44 - padding only - out-of-order + duplicate - dropped as duplicate - // - if err := b.snRangeMap.ExcludeRange(flowState.ExtSequenceNumber, flowState.ExtSequenceNumber+1); err != nil { - b.logger.Errorw( - "could not exclude range", err, - "sn", rtpPacket.SequenceNumber, - "esn", flowState.ExtSequenceNumber, - "rtpStats", b.rtpStats, - "snRangeMap", b.snRangeMap, - ) - } - } - return - } - - if !flowState.IsOutOfOrder && rtpPacket.PayloadType != b.payloadType && b.codecType == webrtc.RTPCodecTypeVideo { - b.handleCodecChange(rtpPacket.PayloadType) - } - - // add to RTX buffer using sequence number after accounting for dropped padding only packets - snAdjustment, err := b.snRangeMap.GetValue(flowState.ExtSequenceNumber) - if err != nil { - b.logger.Errorw( - "could not get sequence number adjustment", err, - "sn", rtpPacket.SequenceNumber, - "esn", flowState.ExtSequenceNumber, - "payloadSize", len(rtpPacket.Payload), - "rtpStats", b.rtpStats, - "snRangeMap", b.snRangeMap, - ) - return - } - flowState.ExtSequenceNumber -= snAdjustment - rtpPacket.Header.SequenceNumber = uint16(flowState.ExtSequenceNumber) - _, err = b.bucket.AddPacketWithSequenceNumber(rawPkt, flowState.ExtSequenceNumber) - if err != nil { - if !flowState.IsDuplicate { - if errors.Is(err, bucket.ErrPacketTooOld) { - packetTooOldCount := b.packetTooOldCount.Inc() - if (packetTooOldCount-1)%100 == 0 { - b.logger.Warnw( - "could not add packet to bucket", err, - "count", packetTooOldCount, - "flowState", &flowState, - "snAdjustment", snAdjustment, - "incomingSequenceNumber", flowState.ExtSequenceNumber+snAdjustment, - "rtpStats", b.rtpStats, - "snRangeMap", b.snRangeMap, - ) - } - } else if err != bucket.ErrRTXPacket { - b.logger.Warnw( - "could not add packet to bucket", err, - "flowState", &flowState, - "snAdjustment", snAdjustment, - "incomingSequenceNumber", flowState.ExtSequenceNumber+snAdjustment, - "rtpStats", b.rtpStats, - "snRangeMap", b.snRangeMap, - ) - } - } - return - } - - ep := b.getExtPacket(rtpPacket, arrivalTime, isBuffered, isRestart, flowState) - if ep == nil { - return - } - b.extPackets.PushBack(ep) - - if b.extPackets.Len() > b.bucket.Capacity() { - if (b.extPacketTooMuchCount.Inc()-1)%100 == 0 { - b.logger.Warnw("too much ext packets", nil, "count", b.extPackets.Len()) - } - } - - b.doFpsCalc(ep) -} - -func (b *Buffer) patchExtPacket(ep *ExtPacket, buf []byte) *ExtPacket { - n, err := b.getPacket(buf, ep.ExtSequenceNumber) - if err != nil { - packetNotFoundCount := b.packetNotFoundCount.Inc() - if (packetNotFoundCount-1)%20 == 0 { - b.logger.Warnw( - "could not get packet from bucket", err, - "sn", ep.Packet.SequenceNumber, - "headSN", b.bucket.HeadSequenceNumber(), - "count", packetNotFoundCount, - "rtpStats", b.rtpStats, - "snRangeMap", b.snRangeMap, - ) - } - return nil - } - ep.RawPacket = buf[:n] - - // patch RTP packet to point payload to new buffer - pkt := *ep.Packet - payloadStart := ep.Packet.Header.MarshalSize() - payloadEnd := payloadStart + len(ep.Packet.Payload) - if payloadEnd > n { - b.logger.Warnw("unexpected marshal size", nil, "max", n, "need", payloadEnd) - return nil - } - pkt.Payload = buf[payloadStart:payloadEnd] - ep.Packet = &pkt - - return ep -} - -func (b *Buffer) doFpsCalc(ep *ExtPacket) { - if b.paused || b.frameRateCalculated || len(ep.Packet.Payload) == 0 { - return - } - spatial := ep.Spatial - if spatial < 0 || int(spatial) >= len(b.frameRateCalculator) { - spatial = 0 - } - if fr := b.frameRateCalculator[spatial]; fr != nil { - if fr.RecvPacket(ep) { - complete := true - for _, fr2 := range b.frameRateCalculator { - if fr2 != nil && !fr2.Completed() { - complete = false - break - } - } - if complete { - b.frameRateCalculated = true - if f := b.onFpsChanged; f != nil { - go f() - } - } - } - } -} - -func (b *Buffer) handleCodecChange(newPT uint8) { - var ( - codecFound, rtxFound bool - rtxPt uint8 - newCodec webrtc.RTPCodecParameters - ) - for _, codec := range b.rtpParameters.Codecs { - if !codecFound && uint8(codec.PayloadType) == newPT { - newCodec = codec - codecFound = true - } - - if mime.IsMimeTypeStringRTX(codec.MimeType) && strings.Contains(codec.SDPFmtpLine, fmt.Sprintf("apt=%d", newPT)) { - rtxFound = true - rtxPt = uint8(codec.PayloadType) - } - - if codecFound && rtxFound { - break - } - } - if !codecFound { - b.logger.Errorw("could not find codec for new payload type", nil, "pt", newPT, "rtpParameters", b.rtpParameters) - return - } - b.logger.Infow( - "codec changed", - "oldPayload", b.payloadType, "newPayload", newPT, - "oldRtxPayload", b.rtxPayloadType, "newRtxPayload", rtxPt, - "oldMime", b.mime, "newMime", newCodec.MimeType) - b.payloadType = newPT - b.rtxPayloadType = rtxPt - b.mime = mime.NormalizeMimeType(newCodec.MimeType) - b.frameRateCalculated = false - - if b.ddExtID != 0 { - b.createDDParserAndFrameRateCalculator() - } - - if b.frameRateCalculator[0] == nil { - b.createFrameRateCalculator() - } - - b.bucket.ResyncOnNextPacket() - - if f := b.onCodecChange; f != nil { - go f(newCodec) - } - - if mime.IsMimeTypeVideo(b.mime) { - go b.seedKeyFrame(b.keyFrameSeederGeneration.Inc()) - } -} - -func (b *Buffer) updateStreamState(p *rtp.Packet, arrivalTime int64) rtpstats.RTPFlowState { - flowState := b.rtpStats.Update( +func (b *Buffer) calc(rawPkt []byte, rtpPacket *rtp.Packet, arrivalTime int64, isBuffered bool, isRTX bool) { + b.BufferBase.HandleIncomingPacketLocked( + rawPkt, + rtpPacket, arrivalTime, - p.Header.SequenceNumber, - p.Header.Timestamp, - p.Header.Marker, - p.Header.MarshalSize(), - len(p.Payload), - int(p.PaddingSize), + isBuffered, + isRTX, + nil, + 0, ) - if b.nacker != nil { - b.nacker.Remove(p.SequenceNumber) + b.doNACKs() - for lost := flowState.LossStartInclusive; lost != flowState.LossEndExclusive; lost++ { - b.nacker.Push(uint16(lost)) - } - } - - return flowState -} - -func (b *Buffer) processHeaderExtensions(p *rtp.Packet, arrivalTime int64, isRTX bool) { - if b.audioLevelExtID != 0 && !isRTX { - if !b.latestTSForAudioLevelInitialized { - b.latestTSForAudioLevelInitialized = true - b.latestTSForAudioLevel = p.Timestamp - } - if e := p.GetExtension(b.audioLevelExtID); e != nil { - ext := rtp.AudioLevelExtension{} - if err := ext.Unmarshal(e); err == nil { - if (p.Timestamp - b.latestTSForAudioLevel) < (1 << 31) { - duration := (int64(p.Timestamp) - int64(b.latestTSForAudioLevel)) * 1e3 / int64(b.clockRate) - if duration > 0 { - b.audioLevel.Observe(ext.Level, uint32(duration), arrivalTime) - } - - b.latestTSForAudioLevel = p.Timestamp - } - } - } - } -} - -func (b *Buffer) getExtPacket( - rtpPacket *rtp.Packet, - arrivalTime int64, - isBuffered bool, - isRestart bool, - flowState rtpstats.RTPFlowState, -) *ExtPacket { - ep := ExtPacketFactory.Get().(*ExtPacket) - *ep = ExtPacket{ - Arrival: arrivalTime, - ExtSequenceNumber: flowState.ExtSequenceNumber, - ExtTimestamp: flowState.ExtTimestamp, - Packet: rtpPacket, - VideoLayer: VideoLayer{ - Spatial: InvalidLayerSpatial, - Temporal: InvalidLayerTemporal, - }, - IsOutOfOrder: flowState.IsOutOfOrder, - IsBuffered: isBuffered, - IsRestart: isRestart, - } - - if len(rtpPacket.Payload) == 0 { - // padding only packet, nothing else to do - return ep - } - - ep.Temporal = 0 - var videoSize []VideoSize - if b.ddParser != nil { - ddVal, videoLayer, err := b.ddParser.Parse(ep.Packet) - if err != nil { - if errors.Is(err, ErrDDExtentionNotFound) { - if b.mime == mime.MimeTypeVP8 || b.mime == mime.MimeTypeVP9 { - b.logger.Infow("dd extension not found, disable dd parser") - b.ddParser = nil - b.createFrameRateCalculator() - } - } else { - ReleaseExtPacket(ep) - return nil - } - } else if ddVal != nil { - ep.DependencyDescriptor = ddVal - ep.VideoLayer = videoLayer - videoSize = ExtractDependencyDescriptorVideoSize(ddVal.Descriptor) - // DD-TODO : notify active decode target change if changed. - } - } - - switch b.mime { - case mime.MimeTypeVP8: - vp8Packet := VP8{} - if err := vp8Packet.Unmarshal(rtpPacket.Payload); err != nil { - b.logger.Warnw("could not unmarshal VP8 packet", err) - ReleaseExtPacket(ep) - return nil - } - ep.KeyFrame = vp8Packet.IsKeyFrame - if ep.DependencyDescriptor == nil { - ep.Temporal = int32(vp8Packet.TID) - - if ep.KeyFrame { - if sz := ExtractVP8VideoSize(&vp8Packet, rtpPacket.Payload); sz.Width > 0 && sz.Height > 0 { - videoSize = append(videoSize, sz) - } - } - } else { - // vp8 with DependencyDescriptor enabled, use the TID from the descriptor - vp8Packet.TID = uint8(ep.Temporal) - } - ep.Payload = vp8Packet - ep.Spatial = InvalidLayerSpatial // vp8 don't have spatial scalability, reset to invalid - - case mime.MimeTypeVP9: - if ep.DependencyDescriptor == nil { - var vp9Packet codecs.VP9Packet - _, err := vp9Packet.Unmarshal(rtpPacket.Payload) - if err != nil { - b.logger.Warnw("could not unmarshal VP9 packet", err) - ReleaseExtPacket(ep) - return nil - } - ep.VideoLayer = VideoLayer{ - Spatial: int32(vp9Packet.SID), - Temporal: int32(vp9Packet.TID), - } - ep.Payload = vp9Packet - ep.KeyFrame = IsVP9KeyFrame(&vp9Packet, rtpPacket.Payload) - - if ep.KeyFrame { - for i := 0; i < len(vp9Packet.Width); i++ { - videoSize = append(videoSize, VideoSize{ - Width: uint32(vp9Packet.Width[i]), - Height: uint32(vp9Packet.Height[i]), - }) - } - } - } else { - ep.KeyFrame = IsVP9KeyFrame(nil, rtpPacket.Payload) - } - - case mime.MimeTypeH264: - ep.KeyFrame = IsH264KeyFrame(rtpPacket.Payload) - ep.Spatial = InvalidLayerSpatial // h.264 don't have spatial scalability, reset to invalid - - // Check H264 key frame video size - if ep.KeyFrame { - if sz := ExtractH264VideoSize(rtpPacket.Payload); sz.Width > 0 && sz.Height > 0 { - videoSize = append(videoSize, sz) - } - } - - case mime.MimeTypeAV1: - ep.KeyFrame = IsAV1KeyFrame(rtpPacket.Payload) - - case mime.MimeTypeH265: - ep.KeyFrame = IsH265KeyFrame(rtpPacket.Payload) - if ep.DependencyDescriptor == nil { - if len(rtpPacket.Payload) < 2 { - b.logger.Warnw("invalid H265 packet", nil) - ReleaseExtPacket(ep) - return nil - } - ep.VideoLayer = VideoLayer{ - Temporal: int32(rtpPacket.Payload[1]&0x07) - 1, - } - ep.Spatial = InvalidLayerSpatial - - if ep.KeyFrame { - if sz := ExtractH265VideoSize(rtpPacket.Payload); sz.Width > 0 && sz.Height > 0 { - videoSize = append(videoSize, sz) - } - } - } - } - - if ep.KeyFrame { - if b.rtpStats != nil { - b.rtpStats.UpdateKeyFrame(1) - } - } - - if b.absCaptureTimeExtID != 0 { - extData := rtpPacket.GetExtension(b.absCaptureTimeExtID) - - var actExt act.AbsCaptureTime - if err := actExt.Unmarshal(extData); err == nil { - ep.AbsCaptureTimeExt = &actExt - } - } - - if len(videoSize) > 0 { - b.checkVideoSizeChange(videoSize) - } - - return ep -} - -func (b *Buffer) flushExtPackets() { - b.Lock() - defer b.Unlock() - b.flushExtPacketsLocked() -} - -func (b *Buffer) flushExtPacketsLocked() { - for b.extPackets.Len() > 0 { - ep := b.extPackets.PopFront() - ReleaseExtPacket(ep) - } - b.extPackets.Clear() + b.doReports(arrivalTime) } func (b *Buffer) doNACKs() { - if b.nacker == nil { - return - } - - if r, numSeqNumsNacked := b.buildNACKPacket(); r != nil { + if r := b.buildNACKPacket(); r != nil { if cb := b.onRtcpFeedback; cb != nil { cb(r) } - if b.rtpStats != nil { - b.rtpStats.UpdateNack(uint32(numSeqNumsNacked)) - } } } +func (b *Buffer) buildNACKPacket() []rtcp.Packet { + if nacks := b.BufferBase.GetNACKPairsLocked(); len(nacks) > 0 { + ssrc := b.BufferBase.SSRC() + pkts := []rtcp.Packet{&rtcp.TransportLayerNack{ + SenderSSRC: ssrc, + MediaSSRC: ssrc, + Nacks: nacks, + }} + return pkts + } + return nil +} + func (b *Buffer) doReports(arrivalTime int64) { - if arrivalTime-b.lastReport < ReportDelta { + if arrivalTime-b.lastReportAt < rtcpReceiverReportDelta { return } - - b.lastReport = arrivalTime + b.lastReportAt = arrivalTime // RTCP reports pkts := b.getRTCP() @@ -1158,105 +373,6 @@ func (b *Buffer) doReports(arrivalTime int64) { cb(pkts) } } - - b.mayGrowBucket() -} - -func (b *Buffer) mayGrowBucket() { - cap := b.bucket.Capacity() - maxPkts := b.maxVideoPkts - if b.codecType == webrtc.RTPCodecTypeAudio { - maxPkts = b.maxAudioPkts - } - if cap >= maxPkts { - return - } - oldCap := cap - if deltaInfo := b.rtpStats.DeltaInfo(b.ppsSnapshotId); deltaInfo != nil { - duration := deltaInfo.EndTime.Sub(deltaInfo.StartTime) - if duration > 500*time.Millisecond { - pps := int(time.Duration(deltaInfo.Packets) * time.Second / duration) - for pps > cap && cap < maxPkts { - cap = b.bucket.Grow() - } - if cap > oldCap { - b.logger.Infow( - "grow bucket", - "from", oldCap, - "to", cap, - "pps", pps, - "deltaInfo", deltaInfo, - "rtpStats", b.rtpStats, - ) - } - } - } -} - -func (b *Buffer) buildNACKPacket() ([]rtcp.Packet, int) { - if nacks, numSeqNumsNacked := b.nacker.Pairs(); len(nacks) > 0 { - pkts := []rtcp.Packet{&rtcp.TransportLayerNack{ - SenderSSRC: b.mediaSSRC, - MediaSSRC: b.mediaSSRC, - Nacks: nacks, - }} - return pkts, numSeqNumsNacked - } - return nil, 0 -} - -func (b *Buffer) buildReceptionReport() *rtcp.ReceptionReport { - if b.rtpStats == nil { - return nil - } - - proxyLoss := b.lastFractionLostToReport - if b.codecType == webrtc.RTPCodecTypeAudio && !b.enableAudioLossProxying { - proxyLoss = 0 - } - - return b.rtpStats.GetRtcpReceptionReport(b.mediaSSRC, proxyLoss, b.rrSnapshotId) -} - -func (b *Buffer) SetSenderReportData(rtpTime uint32, ntpTime uint64, packets uint32, octets uint32) { - b.RLock() - srData := &livekit.RTCPSenderReportState{ - RtpTimestamp: rtpTime, - NtpTimestamp: ntpTime, - At: mono.UnixNano(), - Packets: packets, - Octets: uint64(octets), - } - - didSet := false - if b.rtpStats != nil { - didSet = b.rtpStats.SetRtcpSenderReportData(srData) - } - b.RUnlock() - - if didSet { - if cb := b.getOnRtcpSenderReport(); cb != nil { - cb() - } - } -} - -func (b *Buffer) GetSenderReportData() *livekit.RTCPSenderReportState { - b.RLock() - defer b.RUnlock() - - if b.rtpStats != nil { - return b.rtpStats.GetRtcpSenderReportData() - } - - return nil -} - -func (b *Buffer) SetLastFractionLostReport(lost uint8) { - b.Lock() - defer b.Unlock() - - b.lastFractionLostToReport = lost } func (b *Buffer) getRTCP() []rtcp.Packet { @@ -1265,7 +381,7 @@ func (b *Buffer) getRTCP() []rtcp.Packet { rr := b.buildReceptionReport() if rr != nil { pkts = append(pkts, &rtcp.ReceiverReport{ - SSRC: b.mediaSSRC, + SSRC: b.BufferBase.SSRC(), Reports: []rtcp.ReceptionReport{*rr}, }) } @@ -1273,18 +389,20 @@ func (b *Buffer) getRTCP() []rtcp.Packet { return pkts } -func (b *Buffer) GetPacket(buff []byte, esn uint64) (int, error) { +func (b *Buffer) buildReceptionReport() *rtcp.ReceptionReport { + proxyLoss := b.lastFractionLostToReport + if b.codecType == webrtc.RTPCodecTypeAudio && !b.enableAudioLossProxying { + proxyLoss = 0 + } + + return b.BufferBase.GetRtcpReceptionReportLocked(proxyLoss) +} + +func (b *Buffer) SetLastFractionLostReport(lost uint8) { b.Lock() defer b.Unlock() - return b.getPacket(buff, esn) -} - -func (b *Buffer) getPacket(buff []byte, esn uint64) (int, error) { - if b.closed.Load() { - return 0, io.EOF - } - return b.bucket.GetPacket(buff, esn) + b.lastFractionLostToReport = lost } func (b *Buffer) OnRtcpFeedback(fn func(fb []rtcp.Packet)) { @@ -1300,19 +418,6 @@ func (b *Buffer) getOnRtcpFeedback() func(fb []rtcp.Packet) { return b.onRtcpFeedback } -func (b *Buffer) OnRtcpSenderReport(fn func()) { - b.Lock() - b.onRtcpSenderReport = fn - b.Unlock() -} - -func (b *Buffer) getOnRtcpSenderReport() func() { - b.RLock() - defer b.RUnlock() - - return b.onRtcpSenderReport -} - func (b *Buffer) OnFinalRtpStats(fn func(*livekit.RTPStats)) { b.Lock() b.onFinalRtpStats = fn @@ -1325,184 +430,3 @@ func (b *Buffer) getOnFinalRtpStats() func(*livekit.RTPStats) { return b.onFinalRtpStats } - -// GetMediaSSRC returns the associated SSRC of the RTP stream -func (b *Buffer) GetMediaSSRC() uint32 { - return b.mediaSSRC -} - -// GetClockRate returns the RTP clock rate -func (b *Buffer) GetClockRate() uint32 { - return b.clockRate -} - -func (b *Buffer) GetStats() *livekit.RTPStats { - b.RLock() - defer b.RUnlock() - - if b.rtpStats == nil { - return nil - } - - return b.rtpStats.ToProto() -} - -func (b *Buffer) GetDeltaStats() *StreamStatsWithLayers { - b.RLock() - defer b.RUnlock() - - if b.rtpStats == nil { - return nil - } - - deltaStats := b.rtpStats.DeltaInfo(b.deltaStatsSnapshotId) - if deltaStats == nil { - return nil - } - - return &StreamStatsWithLayers{ - RTPStats: deltaStats, - Layers: map[int32]*rtpstats.RTPDeltaInfo{ - 0: deltaStats, - }, - } -} - -func (b *Buffer) GetLastSenderReportTime() time.Time { - b.RLock() - defer b.RUnlock() - - if b.rtpStats == nil { - return time.Time{} - } - - return b.rtpStats.LastSenderReportTime() -} - -func (b *Buffer) GetAudioLevel() (float64, bool) { - b.RLock() - defer b.RUnlock() - - if b.audioLevel == nil { - return 0, false - } - - return b.audioLevel.GetLevel(mono.UnixNano()) -} - -func (b *Buffer) OnFpsChanged(f func()) { - b.Lock() - b.onFpsChanged = f - b.Unlock() -} - -func (b *Buffer) OnVideoSizeChanged(fn func([]VideoSize)) { - b.Lock() - b.onVideoSizeChanged = fn - b.Unlock() -} - -// checkVideoSizeChange checks if video size has changed for a specific spatial layer and fires callback -func (b *Buffer) checkVideoSizeChange(videoSizes []VideoSize) { - if len(videoSizes) > len(b.currentVideoSize) { - b.logger.Warnw("video size index out of range", nil, "newSize", videoSizes, "currentVideoSize", b.currentVideoSize) - return - } - - if len(videoSizes) < len(b.currentVideoSize) { - videoSizes = append(videoSizes, make([]VideoSize, len(b.currentVideoSize)-len(videoSizes))...) - } - - changed := false - for i, sz := range videoSizes { - if b.currentVideoSize[i].Width != sz.Width || b.currentVideoSize[i].Height != sz.Height { - changed = true - break - } - } - - if changed { - b.logger.Debugw("video size changed", "from", b.currentVideoSize, "to", videoSizes) - copy(b.currentVideoSize[:], videoSizes[:]) - if b.onVideoSizeChanged != nil { - go b.onVideoSizeChanged(videoSizes) - } - } -} - -func (b *Buffer) GetTemporalLayerFpsForSpatial(layer int32) []float32 { - if int(layer) >= len(b.frameRateCalculator) { - return nil - } - - if fc := b.frameRateCalculator[layer]; fc != nil { - return fc.GetFrameRate() - } - return nil -} - -func (b *Buffer) seedKeyFrame(keyFrameSeederGeneration int32) { - // a key frame is needed especially when using Dependency Descriptor - // to get the DD structure which is used in parsing subsequent packets, - // till then packets are dropped which results in stream tracker not - // getting any data which means it does not declare layer start. - // - // send gratuitous PLIs for some time or until a key frame is seen to - // get the engine rolling - b.logger.Debugw("starting key frame seeder") - timer := time.NewTimer(30 * time.Second) - defer timer.Stop() - - ticker := time.NewTicker(time.Second) - defer ticker.Stop() - - initialCount := uint32(0) - b.RLock() - rtpStats := b.rtpStats - b.RUnlock() - if rtpStats != nil { - initialCount, _ = rtpStats.KeyFrame() - } - - for { - if b.closed.Load() || b.keyFrameSeederGeneration.Load() != keyFrameSeederGeneration { - b.logger.Debugw("stopping key frame seeder: stopped") - return - } - - select { - case <-timer.C: - b.logger.Debugw("stopping key frame seeder: timeout") - return - - case <-ticker.C: - if rtpStats != nil { - cnt, last := rtpStats.KeyFrame() - if cnt > initialCount { - b.logger.Debugw( - "stopping key frame seeder: received key frame", - "keyFrameCountInitial", initialCount, - "keyFrameCount", cnt, - "lastKeyFrame", last, - ) - return - } - - b.SendPLI(false) - } - } - } -} - -// --------------------------------------------------------------- - -func ReleaseExtPacket(extPkt *ExtPacket) { - if extPkt == nil { - return - } - - ReleaseExtDependencyDescriptor(extPkt.DependencyDescriptor) - - *extPkt = ExtPacket{} - ExtPacketFactory.Put(extPkt) -} diff --git a/pkg/sfu/buffer/buffer_base.go b/pkg/sfu/buffer/buffer_base.go new file mode 100644 index 000000000..071ed3797 --- /dev/null +++ b/pkg/sfu/buffer/buffer_base.go @@ -0,0 +1,1493 @@ +// 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 buffer + +import ( + "errors" + "fmt" + "io" + "strings" + "sync" + "time" + + "github.com/gammazero/deque" + "github.com/pion/rtcp" + "github.com/pion/rtp" + "github.com/pion/rtp/codecs" + "github.com/pion/sdp/v3" + "github.com/pion/webrtc/v4" + "go.uber.org/atomic" + + "github.com/livekit/livekit-server/pkg/sfu/audio" + "github.com/livekit/livekit-server/pkg/sfu/mime" + act "github.com/livekit/livekit-server/pkg/sfu/rtpextension/abscapturetime" + dd "github.com/livekit/livekit-server/pkg/sfu/rtpextension/dependencydescriptor" + "github.com/livekit/livekit-server/pkg/sfu/rtpstats" + "github.com/livekit/livekit-server/pkg/sfu/utils" + "github.com/livekit/mediatransportutil/pkg/bucket" + "github.com/livekit/mediatransportutil/pkg/nack" + "github.com/livekit/protocol/livekit" + "github.com/livekit/protocol/logger" + "github.com/livekit/protocol/utils/mono" +) + +var ( + ExtPacketFactory = &sync.Pool{ + New: func() any { + return &ExtPacket{} + }, + } +) + +func ReleaseExtPacket(extPkt *ExtPacket) { + if extPkt == nil { + return + } + + ReleaseExtDependencyDescriptor(extPkt.DependencyDescriptor) + + *extPkt = ExtPacket{} + ExtPacketFactory.Put(extPkt) +} + +// -------------------------------------- + +type ExtPacket struct { + VideoLayer + Arrival int64 + ExtSequenceNumber uint64 + ExtTimestamp uint64 + Packet *rtp.Packet + Payload any + KeyFrame bool + RawPacket []byte + DependencyDescriptor *ExtDependencyDescriptor + AbsCaptureTimeExt *act.AbsCaptureTime + IsOutOfOrder bool + IsBuffered bool + IsRestart bool +} + +// VideoSize represents video resolution +type VideoSize struct { + Width uint32 + Height uint32 +} + +type BufferProvider interface { + SetLogger(lgr logger.Logger) + SetAudioLevelParams(audioLevelParams audio.AudioLevelParams) + SetStreamRestartDetection(enable bool) + SetPLIThrottle(duration int64) + SetRTT(rtt uint32) + SetPaused(paused bool) + + SendPLI(force bool) + + ReadExtended(buf []byte) (*ExtPacket, error) + GetPacket(buf []byte, esn uint64) (int, error) + + GetAudioLevel() (float64, bool) + GetTemporalLayerFpsForSpatial(layer int32) []float32 + GetStats() *livekit.RTPStats + GetDeltaStats() *StreamStatsWithLayers + GetDeltaStatsLite() *rtpstats.RTPDeltaInfoLite + GetLastSenderReportTime() time.Time + GetNACKPairs() []rtcp.NackPair + + SetSenderReportData(srData *livekit.RTCPSenderReportState) + GetSenderReportData() *livekit.RTCPSenderReportState + + OnRtcpSenderReport(fn func()) + OnFpsChanged(f func()) + OnVideoSizeChanged(fn func([]VideoSize)) + OnCodecChange(fn func(webrtc.RTPCodecParameters)) + + StartKeyFrameSeeder() + StopKeyFrameSeeder() + + HandleIncomingPacket( + rawPkt []byte, + rtpPacket *rtp.Packet, + arrivalTime int64, + isBuffered bool, + isRTX bool, + skippedSeqs []uint16, + oobSequenceNumber uint16, + ) (*ExtPacket, error) + + CloseWithReason(reason string) (*livekit.RTPStats, error) +} + +const ( + bucketCapCheckInterval = 1e9 +) + +type BufferBaseParams struct { + SSRC uint32 + MaxVideoPkts int + MaxAudioPkts int + LoggerComponents []string + SendPLI func() + IsReportingEnabled bool + IsOOBSequenceNumber bool +} + +type BufferBase struct { + sync.RWMutex + + params BufferBaseParams + + readCond *sync.Cond + + bucket *bucket.Bucket[uint64, uint16] + lastBucketCapCheckAt int64 + + nacker *nack.NackQueue + rtpStatsLite *rtpstats.RTPStatsReceiverLite + liteStatsSnapshotId uint32 + + extPackets deque.Deque[*ExtPacket] + + codecType webrtc.RTPCodecType + closeOnce sync.Once + clockRate uint32 + mime mime.MimeType + + rtpParameters webrtc.RTPParameters + payloadType uint8 + rtxPayloadType uint8 + + snRangeMap *utils.RangeMap[uint64, uint64] + + audioLevelParams audio.AudioLevelParams + audioLevel *audio.AudioLevel + audioLevelExtID uint8 + latestTSForAudioLevelInitialized bool + latestTSForAudioLevel uint32 + + enableStreamRestartDetection bool + + pliThrottle int64 + + rtpStats *rtpstats.RTPStatsReceiver + ppsSnapshotId uint32 + rrSnapshotId uint32 + deltaStatsSnapshotId uint32 + + // callbacks + onRtcpSenderReport func() + onFpsChanged func() + onVideoSizeChanged func([]VideoSize) + onCodecChange func(webrtc.RTPCodecParameters) + + // video size tracking for multiple spatial layers + currentVideoSize [DefaultMaxLayerSpatial + 1]VideoSize + + logger logger.Logger + + // dependency descriptor + ddExtID uint8 + ddParser *DependencyDescriptorParser + + isPaused bool + frameRateCalculator [DefaultMaxLayerSpatial + 1]FrameRateCalculator + frameRateCalculated bool + + packetNotFoundCount atomic.Uint32 + packetTooOldCount atomic.Uint32 + extPacketTooMuchCount atomic.Uint32 + + absCaptureTimeExtID uint8 + + keyFrameSeederGeneration atomic.Int32 + + isClosed atomic.Bool +} + +func NewBufferBase(params BufferBaseParams) *BufferBase { + l := logger.GetLogger() // will be reset with correct context via SetLogger + for _, component := range params.LoggerComponents { + l = l.WithComponent(component) + } + l = l.WithValues("ssrc", params.SSRC) + + b := &BufferBase{ + params: params, + lastBucketCapCheckAt: mono.UnixNano(), + snRangeMap: utils.NewRangeMap[uint64, uint64](100), + pliThrottle: int64(500 * time.Millisecond), + logger: l, + } + b.readCond = sync.NewCond(&b.RWMutex) + b.extPackets.SetBaseCap(128) + return b +} + +func (b *BufferBase) SSRC() uint32 { + return b.params.SSRC +} + +func (b *BufferBase) MaxVideoPkts() int { + return b.params.MaxVideoPkts +} + +func (b *BufferBase) MaxAudioPkts() int { + return b.params.MaxAudioPkts +} + +func (b *BufferBase) SetLogger(lgr logger.Logger) { + b.Lock() + defer b.Unlock() + + for _, component := range b.params.LoggerComponents { + lgr = lgr.WithComponent(component) + } + lgr = lgr.WithValues("ssrc", b.params.SSRC) + b.logger = lgr + + if b.rtpStats != nil { + b.rtpStats.SetLogger(b.logger) + } + + if b.rtpStatsLite != nil { + b.rtpStatsLite.SetLogger(b.logger) + } +} + +func (b *BufferBase) Bind(rtpParameters webrtc.RTPParameters, codec webrtc.RTPCodecCapability, bitrate int) error { + b.Lock() + defer b.Unlock() + + return b.BindLocked(rtpParameters, codec, bitrate) +} + +func (b *BufferBase) BindLocked(rtpParameters webrtc.RTPParameters, codec webrtc.RTPCodecCapability, bitrate int) error { + b.logger.Debugw("binding track") + if codec.ClockRate == 0 { + b.logger.Warnw("invalid codec", nil, "rtpParameters", rtpParameters, "codec", codec, "bitrate", bitrate) + return errInvalidCodec + } + + b.setupRTPStats(codec.ClockRate) + + b.clockRate = codec.ClockRate + b.mime = mime.NormalizeMimeType(codec.MimeType) + b.rtpParameters = rtpParameters + for _, codecParameter := range rtpParameters.Codecs { + if mime.IsMimeTypeStringEqual(codecParameter.MimeType, codec.MimeType) { + b.payloadType = uint8(codecParameter.PayloadType) + break + } + } + + if b.payloadType == 0 && !mime.IsMimeTypeStringEqual(codec.MimeType, webrtc.MimeTypePCMU) { + b.logger.Warnw( + "could not find payload type for codec", nil, + "codec", codec.MimeType, + "rtpParameters", rtpParameters, + ) + b.payloadType = uint8(rtpParameters.Codecs[0].PayloadType) + } + + // find RTX payload type + for _, codec := range rtpParameters.Codecs { + if mime.IsMimeTypeStringRTX(codec.MimeType) && strings.Contains(codec.SDPFmtpLine, fmt.Sprintf("apt=%d", b.payloadType)) { + b.rtxPayloadType = uint8(codec.PayloadType) + break + } + } + + for _, ext := range rtpParameters.HeaderExtensions { + switch ext.URI { + case dd.ExtensionURI: + if b.ddExtID != 0 { + b.logger.Warnw( + "multiple dependency descriptor extensions found", nil, + "id", ext.ID, + "previous", b.ddExtID, + ) + continue + } + b.ddExtID = uint8(ext.ID) + b.createDDParserAndFrameRateCalculator() + + case sdp.AudioLevelURI: + b.audioLevelExtID = uint8(ext.ID) + b.audioLevel = audio.NewAudioLevel(b.audioLevelParams) + + case act.AbsCaptureTimeURI: + b.absCaptureTimeExtID = uint8(ext.ID) + } + } + + switch { + case mime.IsMimeTypeAudio(b.mime): + b.codecType = webrtc.RTPCodecTypeAudio + b.bucket = bucket.NewBucket[uint64, uint16]( + InitPacketBufferSizeAudio, + bucket.RTPMaxPktSize, + bucket.RTPSeqNumOffset, + ) + + case mime.IsMimeTypeVideo(b.mime): + b.codecType = webrtc.RTPCodecTypeVideo + b.bucket = bucket.NewBucket[uint64, uint16]( + InitPacketBufferSizeVideo, + bucket.RTPMaxPktSize, + bucket.RTPSeqNumOffset, + ) + + if b.frameRateCalculator[0] == nil { + b.createFrameRateCalculator() + } + + if bitrate > 0 { + pps := bitrate / 8 / 1200 + for pps > b.bucket.Capacity() { + if b.bucket.Grow() >= b.params.MaxVideoPkts { + break + } + } + } + + default: + b.codecType = webrtc.RTPCodecType(0) + } + + for _, fb := range codec.RTCPFeedback { + switch fb.Type { + case webrtc.TypeRTCPFBGoogREMB: + b.logger.Debugw("Setting feedback", "type", webrtc.TypeRTCPFBGoogREMB) + b.logger.Debugw("REMB not supported, RTCP feedback will not be generated") + + case webrtc.TypeRTCPFBNACK: + // pion uses a single mediaengine to manage negotiated codecs of peerconnection, that means we can't have different + // codec settings at track level for same codec type, so enable nack for all audio receivers but don't create nack queue + // for red codec. + if b.mime == mime.MimeTypeRED { + break + } + + b.logger.Debugw("Setting feedback", "type", webrtc.TypeRTCPFBNACK) + b.nacker = nack.NewNACKQueue(nack.NackQueueParamsDefault) + } + } + + b.StartKeyFrameSeeder() + + return nil +} + +func (b *BufferBase) CloseWithReason(reason string) (stats *livekit.RTPStats, err error) { + b.closeOnce.Do(func() { + b.isClosed.Store(true) + + b.StopKeyFrameSeeder() + + b.RLock() + rtpStats := b.rtpStats + rtpStatsLite := b.rtpStatsLite + b.readCond.Broadcast() + b.RUnlock() + + if rtpStats != nil { + rtpStats.Stop() + stats = rtpStats.ToProto() + } + if rtpStatsLite != nil { + rtpStatsLite.Stop() + } + + b.logger.Debugw( + "rtp stats", + "direction", "upstream", + "stats", rtpStats, + "statsLite", rtpStatsLite, + "reason", reason, + ) + + go b.flushExtPackets() + }) + return +} + +func (b *BufferBase) IsClosed() bool { + return b.isClosed.Load() +} + +func (b *BufferBase) SetPaused(paused bool) { + b.Lock() + defer b.Unlock() + + b.isPaused = paused +} + +func (b *BufferBase) SetAudioLevelParams(audioLevelParams audio.AudioLevelParams) { + b.Lock() + defer b.Unlock() + + b.audioLevelParams = audioLevelParams +} + +func (b *BufferBase) SetStreamRestartDetection(enable bool) { + b.Lock() + defer b.Unlock() + + b.enableStreamRestartDetection = enable +} + +func (b *BufferBase) setupRTPStats(clockRate uint32) { + b.rtpStats = rtpstats.NewRTPStatsReceiver(rtpstats.RTPStatsParams{ + ClockRate: clockRate, + Logger: b.logger, + }) + b.ppsSnapshotId = b.rtpStats.NewSnapshotId() + if b.params.IsReportingEnabled { + b.rrSnapshotId = b.rtpStats.NewSnapshotId() + b.deltaStatsSnapshotId = b.rtpStats.NewSnapshotId() + } + + if b.params.IsOOBSequenceNumber { + b.rtpStatsLite = rtpstats.NewRTPStatsReceiverLite(rtpstats.RTPStatsParams{ + ClockRate: clockRate, + Logger: b.logger, + }) + b.liteStatsSnapshotId = b.rtpStatsLite.NewSnapshotLiteId() + } +} + +func (b *BufferBase) createDDParserAndFrameRateCalculator() { + if mime.IsMimeTypeSVCCapable(b.mime) || b.mime == mime.MimeTypeVP8 { + frc := NewFrameRateCalculatorDD(b.clockRate, b.logger) + for i := range b.frameRateCalculator { + b.frameRateCalculator[i] = frc.GetFrameRateCalculatorForSpatial(int32(i)) + } + b.ddParser = NewDependencyDescriptorParser( + b.ddExtID, + b.logger, + func(spatial, temporal int32) { + frc.SetMaxLayer(spatial, temporal) + }, + false, + ) + } +} + +func (b *BufferBase) createFrameRateCalculator() { + switch b.mime { + case mime.MimeTypeVP8: + b.frameRateCalculator[0] = NewFrameRateCalculatorVP8(b.clockRate, b.logger) + + case mime.MimeTypeVP9: + frc := NewFrameRateCalculatorVP9(b.clockRate, b.logger) + for i := range b.frameRateCalculator { + b.frameRateCalculator[i] = frc.GetFrameRateCalculatorForSpatial(int32(i)) + } + + case mime.MimeTypeH265: + b.frameRateCalculator[0] = NewFrameRateCalculatorH26x(b.clockRate, b.logger) + } +} + +func (b *BufferBase) ReadExtended(buf []byte) (*ExtPacket, error) { + b.Lock() + for { + if b.isClosed.Load() { + b.Unlock() + return nil, io.EOF + } + + if b.extPackets.Len() > 0 { + ep := b.extPackets.PopFront() + patched := b.patchExtPacket(ep, buf) + if patched == nil { + ReleaseExtPacket(ep) + continue + } + + b.Unlock() + return patched, nil + } + + b.readCond.Wait() + } +} + +func (b *BufferBase) SetPLIThrottle(duration int64) { + b.Lock() + defer b.Unlock() + + b.pliThrottle = duration +} + +func (b *BufferBase) SendPLI(force bool) { + b.RLock() + rtpStats := b.rtpStats + pliThrottle := b.pliThrottle + b.RUnlock() + + if (rtpStats == nil && !force) || !rtpStats.CheckAndUpdatePli(pliThrottle, force) { + return + } + + if b.params.SendPLI != nil { + b.params.SendPLI() + } +} + +func (b *BufferBase) SetRTT(rtt uint32) { + b.Lock() + defer b.Unlock() + + if rtt == 0 { + return + } + + if b.nacker != nil { + b.nacker.SetRTT(rtt) + } + + if b.rtpStats != nil { + b.rtpStats.UpdateRtt(rtt) + } +} + +func (b *BufferBase) WaitRead() { + b.readCond.Wait() +} + +func (b *BufferBase) NotifyRead() { + b.readCond.Broadcast() +} + +func (b *BufferBase) HandleIncomingPacket( + rawPkt []byte, + rtpPacket *rtp.Packet, + arrivalTime int64, + isBuffered bool, + isRTX bool, + skippedSeqs []uint16, + oobSequenceNumber uint16, +) (*ExtPacket, error) { + b.Lock() + defer b.Unlock() + + if b.isClosed.Load() { + return nil, io.EOF + } + + return b.HandleIncomingPacketLocked( + rawPkt, + rtpPacket, + arrivalTime, + isBuffered, + isRTX, + skippedSeqs, + oobSequenceNumber, + ) +} + +func (b *BufferBase) HandleIncomingPacketLocked( + rawPkt []byte, + rtpPacket *rtp.Packet, + arrivalTime int64, + isBuffered bool, + isRTX bool, + skippedSeqs []uint16, + oobSequenceNumber uint16, +) (*ExtPacket, error) { + if rtpPacket == nil { + rtpPacket = &rtp.Packet{} + if err := rtpPacket.Unmarshal(rawPkt); err != nil { + b.logger.Errorw("could not unmarshal RTP packet", err) + return nil, err + } + } + + b.processAudioSsrcLevelHeaderExtension(rtpPacket, arrivalTime) + + if len(skippedSeqs) > 0 { + skippedRtpPkt := rtp.Packet{ + Header: rtpPacket.Header, + } + skippedRtpPkt.Marker = false + // Use the current highest timestamp to prevent the case of old sequence number and newer timestamp. + // It is possible that the skipped packet is older. An example sequence + // - Packet 10, skipped 6, 7, 9 -> Packet 8 is unknown at this point + // - Packet 11, skipped 8 -> this would cause sequence number be older, but using timestamp from Packet 11 will make time stamp diff +ve + skippedRtpPkt.Timestamp = b.rtpStats.HighestTimestamp() + for _, sn := range skippedSeqs { + skippedRtpPkt.SequenceNumber = sn + flowState := b.rtpStats.Update( + arrivalTime, + skippedRtpPkt.Header.SequenceNumber, + skippedRtpPkt.Header.Timestamp, + skippedRtpPkt.Header.Marker, + skippedRtpPkt.Header.MarshalSize(), + len(skippedRtpPkt.Payload), + int(skippedRtpPkt.PaddingSize), + ) + if flowState.UnhandledReason == rtpstats.RTPFlowUnhandledReasonNone && !flowState.IsOutOfOrder { + if err := b.snRangeMap.ExcludeRange(flowState.ExtSequenceNumber, flowState.ExtSequenceNumber+1); err != nil { + b.logger.Errorw( + "could not exclude range", err, + "sequenceNumber", sn, + "extSequenceNumber", flowState.ExtSequenceNumber, + "rtpStats", b.rtpStats, + "rtpStatsLite", b.rtpStatsLite, + "snRangeMap", b.snRangeMap, + "skipped", skippedSeqs, + ) + } + } + } + } + + // do not start on an RTX packet + if isRTX && !b.rtpStats.IsActive() { + return nil, errors.New("cannot start on rtx packet") + } + + isRestart := false + flowState := b.rtpStats.Update( + arrivalTime, + rtpPacket.Header.SequenceNumber, + rtpPacket.Header.Timestamp, + rtpPacket.Header.Marker, + rtpPacket.Header.MarshalSize(), + len(rtpPacket.Payload), + int(rtpPacket.PaddingSize), + ) + switch flowState.UnhandledReason { + case rtpstats.RTPFlowUnhandledReasonNone: + case rtpstats.RTPFlowUnhandledReasonRestart: + if !b.enableStreamRestartDetection { + return nil, fmt.Errorf("unhandled reason: %s", flowState.UnhandledReason.String()) + } + + b.StopKeyFrameSeeder() + + b.rtpStats.Stop() + b.logger.Infow("stream restart - rtp stats", b.rtpStats) + + b.snRangeMap = utils.NewRangeMap[uint64, uint64](100) + b.setupRTPStats(b.clockRate) + b.bucket.ResyncOnNextPacket() + if b.nacker != nil { + b.nacker = nack.NewNACKQueue(nack.NackQueueParamsDefault) + } + b.flushExtPacketsLocked() + + flowState = b.rtpStats.Update( + arrivalTime, + rtpPacket.Header.SequenceNumber, + rtpPacket.Header.Timestamp, + rtpPacket.Header.Marker, + rtpPacket.Header.MarshalSize(), + len(rtpPacket.Payload), + int(rtpPacket.PaddingSize), + ) + isRestart = true + default: + return nil, fmt.Errorf("unhandled reason: %s", flowState.UnhandledReason.String()) + } + + if len(rtpPacket.Payload) == 0 && (!flowState.IsOutOfOrder || flowState.IsDuplicate) { + // drop padding only in-order or duplicate packet + if !flowState.IsOutOfOrder { + // in-order packet - increment sequence number offset for subsequent packets + // Example: + // 40 - regular packet - pass through as sequence number 40 + // 41 - missing packet - don't know what it is, could be padding or not + // 42 - padding only packet - in-order - drop - increment sequence number offset to 1 - + // range[0, 42] = 0 offset + // 41 - arrives out of order - get offset 0 from cache - passed through as sequence number 41 + // 43 - regular packet - offset = 1 (running offset) - passes through as sequence number 42 + // 44 - padding only - in order - drop - increment sequence number offset to 2 + // range[0, 42] = 0 offset, range[43, 44] = 1 offset + // 43 - regular packet - out of order + duplicate - offset = 1 from cache - + // adjusted sequence number is 42, will be dropped by RTX buffer AddPacket method as duplicate + // 45 - regular packet - offset = 2 (running offset) - passed through with adjusted sequence number as 43 + // 44 - padding only - out-of-order + duplicate - dropped as duplicate + // + if err := b.snRangeMap.ExcludeRange(flowState.ExtSequenceNumber, flowState.ExtSequenceNumber+1); err != nil { + b.logger.Errorw( + "could not exclude range", err, + "sn", rtpPacket.SequenceNumber, + "esn", flowState.ExtSequenceNumber, + "rtpStats", b.rtpStats, + "snRangeMap", b.snRangeMap, + ) + } + } + return nil, errors.New("padding only packet") + } + + if !flowState.IsOutOfOrder && rtpPacket.PayloadType != b.payloadType && b.codecType == webrtc.RTPCodecTypeVideo { + b.logger.Infow("possible codec change", "oldPT", b.payloadType, "receivedPT", rtpPacket.PayloadType) + b.handleCodecChange(rtpPacket.PayloadType) + } + + // add to RTX buffer using sequence number after accounting for dropped padding only packets + snAdjustment, err := b.snRangeMap.GetValue(flowState.ExtSequenceNumber) + if err != nil { + b.logger.Errorw( + "could not get sequence number adjustment", err, + "sequenceNumber", rtpPacket.SequenceNumber, + "extSequenceNumber", flowState.ExtSequenceNumber, + "timestamp", rtpPacket.Timestamp, + "extTimestamp", flowState.ExtTimestamp, + "payloadSize", len(rtpPacket.Payload), + "paddingSize", rtpPacket.PaddingSize, + "rtpStats", b.rtpStats, + "rtpStatsLite", b.rtpStatsLite, + "snRangeMap", b.snRangeMap, + ) + return nil, err + } + + flowState.ExtSequenceNumber -= snAdjustment + rtpPacket.Header.SequenceNumber = uint16(flowState.ExtSequenceNumber) + if _, err = b.bucket.AddPacketWithSequenceNumber(rawPkt, flowState.ExtSequenceNumber); err != nil { + if !flowState.IsDuplicate { + if errors.Is(err, bucket.ErrPacketTooOld) { + packetTooOldCount := b.packetTooOldCount.Inc() + if (packetTooOldCount-1)%100 == 0 { + b.logger.Warnw( + "could not add packet to bucket", err, + "count", packetTooOldCount, + "flowState", &flowState, + "snAdjustment", snAdjustment, + "incomingSequenceNumber", flowState.ExtSequenceNumber+snAdjustment, + "rtpStats", b.rtpStats, + "rtpStatsLite", b.rtpStatsLite, + "snRangeMap", b.snRangeMap, + "skipped", skippedSeqs, + ) + } + } else if err != bucket.ErrRTXPacket { + b.logger.Warnw( + "could not add packet to bucket", err, + "flowState", &flowState, + "snAdjustment", snAdjustment, + "incomingSequenceNumber", flowState.ExtSequenceNumber+snAdjustment, + "rtpStats", b.rtpStats, + "rtpStatsLite", b.rtpStatsLite, + "snRangeMap", b.snRangeMap, + "skipped", skippedSeqs, + ) + } + } + return nil, err + } + + ep := b.getExtPacket(rtpPacket, arrivalTime, isBuffered, isRestart, flowState) + if ep == nil { + return nil, errors.New("could not get ext packet") + } + b.extPackets.PushBack(ep) + b.readCond.Broadcast() + + if b.extPackets.Len() > b.bucket.Capacity() { + if (b.extPacketTooMuchCount.Inc()-1)%100 == 0 { + b.logger.Warnw("too much ext packets", nil, "count", b.extPackets.Len()) + } + } + + b.maybeGrowBucket(arrivalTime) + + if b.params.IsOOBSequenceNumber { + b.updateOOBNACKState(oobSequenceNumber, arrivalTime, len(rawPkt)) + } else { + b.updateNACKState(rtpPacket.SequenceNumber, flowState) + } + + return ep, nil +} + +func (b *BufferBase) updateNACKState(sequenceNumber uint16, flowState rtpstats.RTPFlowState) { + if b.nacker == nil { + return + } + + b.nacker.Remove(sequenceNumber) + + for lost := flowState.LossStartInclusive; lost != flowState.LossEndExclusive; lost++ { + b.nacker.Push(uint16(lost)) + } +} + +func (b *BufferBase) updateOOBNACKState(sequenceNumber uint16, arrivalTime int64, size int) { + if b.nacker == nil || !b.params.IsOOBSequenceNumber { + return + } + + fsLite := b.rtpStatsLite.Update(arrivalTime, size, sequenceNumber) + if fsLite.IsNotHandled { + return + } + + b.nacker.Remove(sequenceNumber) + + for lost := fsLite.LossStartInclusive; lost != fsLite.LossEndExclusive; lost++ { + b.nacker.Push(uint16(lost)) + } +} + +func (b *BufferBase) processAudioSsrcLevelHeaderExtension(p *rtp.Packet, arrivalTime int64) { + if b.audioLevelExtID == 0 { + return + } + + if !b.latestTSForAudioLevelInitialized { + b.latestTSForAudioLevelInitialized = true + b.latestTSForAudioLevel = p.Timestamp + } + if e := p.GetExtension(b.audioLevelExtID); e != nil { + ext := rtp.AudioLevelExtension{} + if err := ext.Unmarshal(e); err == nil { + if (p.Timestamp - b.latestTSForAudioLevel) < (1 << 31) { + duration := (int64(p.Timestamp) - int64(b.latestTSForAudioLevel)) * 1e3 / int64(b.clockRate) + if duration > 0 { + b.audioLevel.Observe(ext.Level, uint32(duration), arrivalTime) + } + + b.latestTSForAudioLevel = p.Timestamp + } + } + } +} + +func (b *BufferBase) handleCodecChange(newPT uint8) { + var ( + codecFound, rtxFound bool + rtxPt uint8 + newCodec webrtc.RTPCodecParameters + ) + for _, codec := range b.rtpParameters.Codecs { + if !codecFound && uint8(codec.PayloadType) == newPT { + newCodec = codec + codecFound = true + } + + if mime.IsMimeTypeStringRTX(codec.MimeType) && strings.Contains(codec.SDPFmtpLine, fmt.Sprintf("apt=%d", newPT)) { + rtxFound = true + rtxPt = uint8(codec.PayloadType) + } + + if codecFound && rtxFound { + break + } + } + if !codecFound { + b.logger.Errorw( + "could not find codec for new payload type", nil, + "pt", newPT, + "rtpParameters", b.rtpParameters, + ) + return + } + b.logger.Infow( + "codec changed", + "oldPayload", b.payloadType, "newPayload", newPT, + "oldRtxPayload", b.rtxPayloadType, "newRtxPayload", rtxPt, + "oldMime", b.mime, "newMime", newCodec.MimeType, + ) + b.payloadType = newPT + b.rtxPayloadType = rtxPt + b.mime = mime.NormalizeMimeType(newCodec.MimeType) + b.frameRateCalculated = false + + if b.ddExtID != 0 { + b.createDDParserAndFrameRateCalculator() + } + + if b.frameRateCalculator[0] == nil { + b.createFrameRateCalculator() + } + + b.bucket.ResyncOnNextPacket() + + if f := b.onCodecChange; f != nil { + go f(newCodec) + } + + b.StartKeyFrameSeeder() +} + +func (b *BufferBase) getExtPacket( + rtpPacket *rtp.Packet, + arrivalTime int64, + isBuffered bool, + isRestart bool, + flowState rtpstats.RTPFlowState, +) *ExtPacket { + ep := ExtPacketFactory.Get().(*ExtPacket) + *ep = ExtPacket{ + Arrival: arrivalTime, + ExtSequenceNumber: flowState.ExtSequenceNumber, + ExtTimestamp: flowState.ExtTimestamp, + Packet: rtpPacket, + VideoLayer: VideoLayer{ + Spatial: InvalidLayerSpatial, + Temporal: InvalidLayerTemporal, + }, + IsOutOfOrder: flowState.IsOutOfOrder, + IsBuffered: isBuffered, + IsRestart: isRestart, + } + + if len(ep.Packet.Payload) == 0 { + // padding only packet, nothing else to do + return ep + } + + if err := b.processVideoPacket(ep); err != nil { + ReleaseExtPacket(ep) + return nil + } + + if b.absCaptureTimeExtID != 0 { + extData := rtpPacket.GetExtension(b.absCaptureTimeExtID) + + var actExt act.AbsCaptureTime + if err := actExt.Unmarshal(extData); err == nil { + ep.AbsCaptureTimeExt = &actExt + } + } + + return ep +} + +func (b *BufferBase) processVideoPacket(ep *ExtPacket) error { + if b.codecType != webrtc.RTPCodecTypeVideo { + return nil + } + + ep.Temporal = 0 + var videoSize []VideoSize + if b.ddParser != nil { + ddVal, videoLayer, err := b.ddParser.Parse(ep.Packet) + if err != nil { + if errors.Is(err, ErrDDExtentionNotFound) { + if b.mime == mime.MimeTypeVP8 || b.mime == mime.MimeTypeVP9 { + b.logger.Infow("dd extension not found, disable dd parser") + b.ddParser = nil + b.createFrameRateCalculator() + } + } else { + return err + } + } else if ddVal != nil { + ep.DependencyDescriptor = ddVal + ep.VideoLayer = videoLayer + videoSize = ExtractDependencyDescriptorVideoSize(ddVal.Descriptor) + // DD-TODO : notify active decode target change if changed. + } + } + + switch b.mime { + case mime.MimeTypeVP8: + vp8Packet := VP8{} + if err := vp8Packet.Unmarshal(ep.Packet.Payload); err != nil { + b.logger.Warnw("could not unmarshal VP8 packet", err) + return err + } + ep.KeyFrame = vp8Packet.IsKeyFrame + if ep.DependencyDescriptor == nil { + ep.Temporal = int32(vp8Packet.TID) + + if ep.KeyFrame { + if sz := ExtractVP8VideoSize(&vp8Packet, ep.Packet.Payload); sz.Width > 0 && sz.Height > 0 { + videoSize = append(videoSize, sz) + } + } + } else { + // vp8 with DependencyDescriptor enabled, use the TID from the descriptor + vp8Packet.TID = uint8(ep.Temporal) + } + ep.Payload = vp8Packet + ep.Spatial = InvalidLayerSpatial // vp8 don't have spatial scalability, reset to invalid + + case mime.MimeTypeVP9: + if ep.DependencyDescriptor == nil { + var vp9Packet codecs.VP9Packet + _, err := vp9Packet.Unmarshal(ep.Packet.Payload) + if err != nil { + b.logger.Warnw("could not unmarshal VP9 packet", err) + return err + } + ep.VideoLayer = VideoLayer{ + Spatial: int32(vp9Packet.SID), + Temporal: int32(vp9Packet.TID), + } + ep.Payload = vp9Packet + ep.KeyFrame = IsVP9KeyFrame(&vp9Packet, ep.Packet.Payload) + + if ep.KeyFrame { + for i := 0; i < len(vp9Packet.Width); i++ { + videoSize = append(videoSize, VideoSize{ + Width: uint32(vp9Packet.Width[i]), + Height: uint32(vp9Packet.Height[i]), + }) + } + } + } else { + ep.KeyFrame = IsVP9KeyFrame(nil, ep.Packet.Payload) + } + + case mime.MimeTypeH264: + ep.KeyFrame = IsH264KeyFrame(ep.Packet.Payload) + ep.Spatial = InvalidLayerSpatial // h.264 don't have spatial scalability, reset to invalid + + // Check H264 key frame video size + if ep.KeyFrame { + if sz := ExtractH264VideoSize(ep.Packet.Payload); sz.Width > 0 && sz.Height > 0 { + videoSize = append(videoSize, sz) + } + } + + case mime.MimeTypeAV1: + ep.KeyFrame = IsAV1KeyFrame(ep.Packet.Payload) + + case mime.MimeTypeH265: + ep.KeyFrame = IsH265KeyFrame(ep.Packet.Payload) + if ep.DependencyDescriptor == nil { + if len(ep.Packet.Payload) < 2 { + b.logger.Warnw("invalid H265 packet", nil, "payloadLen", len(ep.Packet.Payload)) + return errors.New("invalid H265 packet") + } + ep.VideoLayer = VideoLayer{ + Temporal: int32(ep.Packet.Payload[1]&0x07) - 1, + } + ep.Spatial = InvalidLayerSpatial + + if ep.KeyFrame { + if sz := ExtractH265VideoSize(ep.Packet.Payload); sz.Width > 0 && sz.Height > 0 { + videoSize = append(videoSize, sz) + } + } + } + } + + if ep.KeyFrame { + if b.rtpStats != nil { + b.rtpStats.UpdateKeyFrame(1) + } + } + + if len(videoSize) > 0 { + b.checkVideoSizeChange(videoSize) + } + + b.doFpsCalc(ep) + + return nil +} + +func (b *BufferBase) patchExtPacket(ep *ExtPacket, buf []byte) *ExtPacket { + n, err := b.getPacketLocked(buf, ep.ExtSequenceNumber) + if err != nil { + packetNotFoundCount := b.packetNotFoundCount.Inc() + if (packetNotFoundCount-1)%20 == 0 { + b.logger.Warnw( + "could not get packet from bucket", err, + "sn", ep.Packet.SequenceNumber, + "headSN", b.bucket.HeadSequenceNumber(), + "count", packetNotFoundCount, + "rtpStats", b.rtpStats, + "rtpStatsLite", b.rtpStatsLite, + "snRangeMap", b.snRangeMap, + ) + } + return nil + } + ep.RawPacket = buf[:n] + + // patch RTP packet to point payload to new buffer + pkt := *ep.Packet + payloadStart := ep.Packet.Header.MarshalSize() + payloadEnd := payloadStart + len(ep.Packet.Payload) + if payloadEnd > n { + b.logger.Warnw("unexpected marshal size", nil, "max", n, "need", payloadEnd) + return nil + } + pkt.Payload = buf[payloadStart:payloadEnd] + ep.Packet = &pkt + + return ep +} + +func (b *BufferBase) flushExtPackets() { + b.Lock() + defer b.Unlock() + b.flushExtPacketsLocked() +} + +func (b *BufferBase) flushExtPacketsLocked() { + for b.extPackets.Len() > 0 { + ep := b.extPackets.PopFront() + ReleaseExtPacket(ep) + } + b.extPackets.Clear() +} + +func (b *BufferBase) maybeGrowBucket(now int64) { + if now-b.lastBucketCapCheckAt < bucketCapCheckInterval { + return + } + b.lastBucketCapCheckAt = now + + cap := b.bucket.Capacity() + maxPkts := b.params.MaxVideoPkts + if b.codecType == webrtc.RTPCodecTypeAudio { + maxPkts = b.params.MaxAudioPkts + } + if cap >= maxPkts { + return + } + + oldCap := cap + if deltaInfo := b.rtpStats.DeltaInfo(b.ppsSnapshotId); deltaInfo != nil { + duration := deltaInfo.EndTime.Sub(deltaInfo.StartTime) + if duration > 500*time.Millisecond { + pps := int(time.Duration(deltaInfo.Packets) * time.Second / duration) + for pps > cap && cap < maxPkts { + cap = b.bucket.Grow() + } + if cap > oldCap { + b.logger.Infow( + "grow bucket", + "from", oldCap, + "to", cap, + "pps", pps, + "deltaInfo", deltaInfo, + "rtpStats", b.rtpStats, + ) + } + } + } +} + +func (b *BufferBase) doFpsCalc(ep *ExtPacket) { + if b.isPaused || b.frameRateCalculated || len(ep.Packet.Payload) == 0 { + return + } + + spatial := ep.Spatial + if spatial < 0 || int(spatial) >= len(b.frameRateCalculator) { + spatial = 0 + } + if fr := b.frameRateCalculator[spatial]; fr != nil { + if fr.RecvPacket(ep) { + complete := true + for _, fr2 := range b.frameRateCalculator { + if fr2 != nil && !fr2.Completed() { + complete = false + break + } + } + if complete { + b.frameRateCalculated = true + if f := b.onFpsChanged; f != nil { + go f() + } + } + } + } +} + +func (b *BufferBase) SetSenderReportData(srData *livekit.RTCPSenderReportState) { + srData.At = mono.UnixNano() + b.RLock() + didSet := false + if b.rtpStats != nil { + didSet = b.rtpStats.SetRtcpSenderReportData(srData) + } + b.RUnlock() + + if didSet { + if cb := b.getOnRtcpSenderReport(); cb != nil { + cb() + } + } +} + +func (b *BufferBase) GetSenderReportData() *livekit.RTCPSenderReportState { + b.RLock() + defer b.RUnlock() + + if b.rtpStats != nil { + return b.rtpStats.GetRtcpSenderReportData() + } + + return nil +} + +func (b *BufferBase) GetPacket(buff []byte, esn uint64) (int, error) { + b.Lock() + defer b.Unlock() + + return b.getPacketLocked(buff, esn) +} + +func (b *BufferBase) getPacketLocked(buff []byte, esn uint64) (int, error) { + if b.isClosed.Load() { + return 0, io.EOF + } + return b.bucket.GetPacket(buff, esn) +} + +func (b *BufferBase) GetStats() *livekit.RTPStats { + b.RLock() + defer b.RUnlock() + + if b.rtpStats == nil { + return nil + } + + return b.rtpStats.ToProto() +} + +func (b *BufferBase) GetDeltaStats() *StreamStatsWithLayers { + b.RLock() + defer b.RUnlock() + + if b.rtpStats == nil { + return nil + } + + deltaStats := b.rtpStats.DeltaInfo(b.deltaStatsSnapshotId) + if deltaStats == nil { + return nil + } + + return &StreamStatsWithLayers{ + RTPStats: deltaStats, + Layers: map[int32]*rtpstats.RTPDeltaInfo{ + 0: deltaStats, + }, + } +} + +func (b *BufferBase) GetDeltaStatsLite() *rtpstats.RTPDeltaInfoLite { + b.RLock() + defer b.RUnlock() + + if b.rtpStatsLite == nil { + return nil + } + + return b.rtpStatsLite.DeltaInfoLite(b.liteStatsSnapshotId) +} + +func (b *BufferBase) GetLastSenderReportTime() time.Time { + b.RLock() + defer b.RUnlock() + + if b.rtpStats == nil { + return time.Time{} + } + + return b.rtpStats.LastSenderReportTime() +} + +func (b *BufferBase) GetAudioLevel() (float64, bool) { + b.RLock() + defer b.RUnlock() + + if b.audioLevel == nil { + return 0, false + } + + return b.audioLevel.GetLevel(mono.UnixNano()) +} + +func (b *BufferBase) OnRtcpSenderReport(fn func()) { + b.Lock() + b.onRtcpSenderReport = fn + b.Unlock() +} + +func (b *BufferBase) getOnRtcpSenderReport() func() { + b.RLock() + defer b.RUnlock() + + return b.onRtcpSenderReport +} + +func (b *BufferBase) OnFpsChanged(f func()) { + b.Lock() + b.onFpsChanged = f + b.Unlock() +} + +func (b *BufferBase) OnVideoSizeChanged(fn func([]VideoSize)) { + b.Lock() + b.onVideoSizeChanged = fn + b.Unlock() +} + +func (b *BufferBase) OnCodecChange(fn func(webrtc.RTPCodecParameters)) { + b.Lock() + b.onCodecChange = fn + b.Unlock() +} + +// checkVideoSizeChange checks if video size has changed for a specific spatial layer and fires callback +func (b *BufferBase) checkVideoSizeChange(videoSizes []VideoSize) { + if len(videoSizes) > len(b.currentVideoSize) { + b.logger.Warnw( + "video size index out of range", nil, + "newSize", videoSizes, + "currentVideoSize", b.currentVideoSize, + ) + return + } + + if len(videoSizes) < len(b.currentVideoSize) { + videoSizes = append(videoSizes, make([]VideoSize, len(b.currentVideoSize)-len(videoSizes))...) + } + + changed := false + for i, sz := range videoSizes { + if b.currentVideoSize[i].Width != sz.Width || b.currentVideoSize[i].Height != sz.Height { + changed = true + break + } + } + + if changed { + b.logger.Debugw("video size changed", "from", b.currentVideoSize, "to", videoSizes) + copy(b.currentVideoSize[:], videoSizes[:]) + if b.onVideoSizeChanged != nil { + go b.onVideoSizeChanged(videoSizes) + } + } +} + +func (b *BufferBase) GetTemporalLayerFpsForSpatial(layer int32) []float32 { + b.RLock() + defer b.RUnlock() + + if int(layer) >= len(b.frameRateCalculator) { + return nil + } + + if fc := b.frameRateCalculator[layer]; fc != nil { + return fc.GetFrameRate() + } + return nil +} + +func (b *BufferBase) StartKeyFrameSeeder() { + if b.codecType == webrtc.RTPCodecTypeVideo { + go b.seedKeyFrame(b.keyFrameSeederGeneration.Inc()) + } +} + +func (b *BufferBase) StopKeyFrameSeeder() { + b.keyFrameSeederGeneration.Inc() +} + +func (b *BufferBase) seedKeyFrame(keyFrameSeederGeneration int32) { + // a key frame is needed especially when using Dependency Descriptor + // to get the DD structure which is used in parsing subsequent packets, + // till then packets are dropped which results in stream tracker not + // getting any data which means it does not declare layer start. + // + // send gratuitous PLIs for some time or until a key frame is seen to + // get the engine rolling + b.logger.Debugw("starting key frame seeder") + timer := time.NewTimer(30 * time.Second) + defer timer.Stop() + + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + + initialCount := uint32(0) + b.RLock() + rtpStats := b.rtpStats + b.RUnlock() + if rtpStats == nil { + b.logger.Debugw("cannot do key frame seeding without stats") + return + } + initialCount, _ = rtpStats.KeyFrame() + + for { + if b.isClosed.Load() || b.keyFrameSeederGeneration.Load() != keyFrameSeederGeneration { + b.logger.Debugw("stopping key frame seeder: stopped") + return + } + + select { + case <-timer.C: + b.logger.Debugw("stopping key frame seeder: timeout") + return + + case <-ticker.C: + cnt, last := rtpStats.KeyFrame() + if cnt > initialCount { + b.logger.Debugw( + "stopping key frame seeder: received key frame", + "keyFrameCountInitial", initialCount, + "keyFrameCount", cnt, + "lastKeyFrame", last, + ) + return + } + + b.SendPLI(false) + } + } +} + +func (b *BufferBase) GetNACKPairs() []rtcp.NackPair { + b.RLock() + defer b.RUnlock() + + return b.GetNACKPairsLocked() +} + +func (b *BufferBase) GetNACKPairsLocked() []rtcp.NackPair { + if b.nacker == nil { + return nil + } + + pairs, numSeqNumsNacked := b.nacker.Pairs() + if !b.params.IsOOBSequenceNumber { + if b.rtpStats != nil { + b.rtpStats.UpdateNack(uint32(numSeqNumsNacked)) + } + } else { + if b.rtpStatsLite != nil { + b.rtpStatsLite.UpdateNack(uint32(numSeqNumsNacked)) + } + } + + return pairs +} + +func (b *BufferBase) GetRtcpReceptionReportLocked(proxyLoss uint8) *rtcp.ReceptionReport { + if b.rtpStats == nil { + return nil + } + + return b.rtpStats.GetRtcpReceptionReport(b.params.SSRC, proxyLoss, b.rrSnapshotId) +} + +// --------------------------------------------------------------- diff --git a/pkg/sfu/receiver.go b/pkg/sfu/receiver.go index f459b1032..5a1482e75 100644 --- a/pkg/sfu/receiver.go +++ b/pkg/sfu/receiver.go @@ -15,201 +15,38 @@ package sfu import ( - "errors" - "io" "strings" "sync" "time" "github.com/pion/rtcp" "github.com/pion/webrtc/v4" - "go.uber.org/atomic" - "github.com/livekit/mediatransportutil/pkg/bucket" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" - "github.com/livekit/protocol/utils" - "github.com/livekit/protocol/utils/mono" - "github.com/livekit/livekit-server/pkg/sfu/audio" "github.com/livekit/livekit-server/pkg/sfu/buffer" "github.com/livekit/livekit-server/pkg/sfu/connectionquality" "github.com/livekit/livekit-server/pkg/sfu/mime" "github.com/livekit/livekit-server/pkg/sfu/rtpstats" - "github.com/livekit/livekit-server/pkg/sfu/streamtracker" - sfuutils "github.com/livekit/livekit-server/pkg/sfu/utils" ) -var ( - ErrReceiverClosed = errors.New("receiver closed") - ErrDownTrackAlreadyExist = errors.New("DownTrack already exist") - ErrBufferNotFound = errors.New("buffer not found") - ErrDuplicateLayer = errors.New("duplicate layer") - ErrInvalidLayer = errors.New("invalid layer") -) - -// -------------------------------------- - -type PLIThrottleConfig struct { - LowQuality time.Duration `yaml:"low_quality,omitempty"` - MidQuality time.Duration `yaml:"mid_quality,omitempty"` - HighQuality time.Duration `yaml:"high_quality,omitempty"` -} - -var ( - DefaultPLIThrottleConfig = PLIThrottleConfig{ - LowQuality: 500 * time.Millisecond, - MidQuality: time.Second, - HighQuality: time.Second, - } -) - -// -------------------------------------- - -type AudioConfig struct { - audio.AudioLevelConfig `yaml:",inline"` - - // enable red encoding downtrack for opus only audio up track - ActiveREDEncoding bool `yaml:"active_red_encoding,omitempty"` - // enable proxying weakest subscriber loss to publisher in RTCP Receiver Report - EnableLossProxying bool `yaml:"enable_loss_proxying,omitempty"` -} - -var ( - DefaultAudioConfig = AudioConfig{ - AudioLevelConfig: audio.DefaultAudioLevelConfig, - } -) - -// -------------------------------------- - -type AudioLevelHandle func(level uint8, duration uint32) - -type Bitrates [buffer.DefaultMaxLayerSpatial + 1][buffer.DefaultMaxLayerTemporal + 1]int64 - -type ReceiverCodecState int - -const ( - ReceiverCodecStateNormal ReceiverCodecState = iota - ReceiverCodecStateSuspended - ReceiverCodecStateInvalid -) - -// TrackReceiver defines an interface receive media from remote peer -type TrackReceiver interface { - TrackID() livekit.TrackID - StreamID() string - - // returns the initial codec of the receiver, it is determined by the track's codec - // and will not change if the codec changes during the session (publisher changes codec) - Codec() webrtc.RTPCodecParameters - Mime() mime.MimeType - VideoLayerMode() livekit.VideoLayer_Mode - HeaderExtensions() []webrtc.RTPHeaderExtensionParameter - IsClosed() bool - - ReadRTP(buf []byte, layer uint8, esn uint64) (int, error) - GetLayeredBitrate() ([]int32, Bitrates) - - GetAudioLevel() (float64, bool) - - SendPLI(layer int32, force bool) - - SetUpTrackPaused(paused bool) - SetMaxExpectedSpatialLayer(layer int32) - - AddDownTrack(track TrackSender) error - DeleteDownTrack(participantID livekit.ParticipantID) - GetDownTracks() []TrackSender - - DebugInfo() map[string]any - - TrackInfo() *livekit.TrackInfo - UpdateTrackInfo(ti *livekit.TrackInfo) - - // Get primary receiver if this receiver represents a RED codec; otherwise it will return itself - GetPrimaryReceiverForRed() TrackReceiver - - // Get red receiver for primary codec, used by forward red encodings for opus only codec - GetRedReceiver() TrackReceiver - - GetTemporalLayerFpsForSpatial(layer int32) []float32 - - GetTrackStats() *livekit.RTPStats - - // AddOnReady adds a function to be called when the receiver is ready, the callback - // could be called immediately if the receiver is ready when the callback is added - AddOnReady(func()) - - AddOnCodecStateChange(func(webrtc.RTPCodecParameters, ReceiverCodecState)) - CodecState() ReceiverCodecState - - // VideoSizes returns the video size parsed from rtp packet for each spatial layer. - VideoSizes() []buffer.VideoSize -} - -type REDTransformer interface { - ForwardRTP(pkt *buffer.ExtPacket, spatialLayer int32) int32 - ForwardRTCPSenderReport( - payloadType webrtc.PayloadType, - layer int32, - publisherSRData *livekit.RTCPSenderReportState, - ) - ResyncDownTracks() - OnStreamRestart() - CanClose() bool - Close() -} - var _ TrackReceiver = (*WebRTCReceiver)(nil) // WebRTCReceiver receives a media track type WebRTCReceiver struct { - logger logger.Logger + *ReceiverBase - pliThrottleConfig PLIThrottleConfig - audioConfig AudioConfig - enableRTPStreamRestartDetection bool - - trackID livekit.TrackID - streamID string - kind webrtc.RTPCodecType - receiver *webrtc.RTPReceiver - codec webrtc.RTPCodecParameters - codecState ReceiverCodecState - codecStateLock sync.Mutex - onCodecStateChange []func(webrtc.RTPCodecParameters, ReceiverCodecState) - isRED bool - onCloseHandler func() - closeOnce sync.Once - closed atomic.Bool - trackInfo atomic.Pointer[livekit.TrackInfo] - videoLayerMode livekit.VideoLayer_Mode + receiver *webrtc.RTPReceiver + onCloseHandler func() onRTCP func([]rtcp.Packet) - bufferMu sync.RWMutex - buffers [buffer.DefaultMaxLayerSpatial + 1]*buffer.Buffer - upTracks [buffer.DefaultMaxLayerSpatial + 1]TrackRemote - videoSizeMu sync.RWMutex - videoSizes [buffer.DefaultMaxLayerSpatial + 1]buffer.VideoSize - onVideoSizeChanged func() - rtt uint32 - - lbThreshold int - - streamTrackerManager *StreamTrackerManager - - downTrackSpreader *sfuutils.DownTrackSpreader[TrackSender] + upTracksMu sync.Mutex + upTracks [buffer.DefaultMaxLayerSpatial + 1]TrackRemote connectionStats *connectionquality.ConnectionStats - - onStatsUpdate func(w *WebRTCReceiver, stat *livekit.AnalyticsStat) - onMaxLayerChange func(mimeType mime.MimeType, maxLayer int32) - - redTransformer atomic.Value // redTransformer interface - - forwardStats *ForwardStats + onStatsUpdate func(w *WebRTCReceiver, stat *livekit.AnalyticsStat) } type ReceiverOpts func(w *WebRTCReceiver) *WebRTCReceiver @@ -217,7 +54,7 @@ type ReceiverOpts func(w *WebRTCReceiver) *WebRTCReceiver // WithPliThrottleConfig indicates minimum time(ms) between sending PLIs func WithPliThrottleConfig(pliThrottleConfig PLIThrottleConfig) ReceiverOpts { return func(w *WebRTCReceiver) *WebRTCReceiver { - w.pliThrottleConfig = pliThrottleConfig + w.ReceiverBase.SetPLIThrottleConfig(pliThrottleConfig) return w } } @@ -225,14 +62,14 @@ func WithPliThrottleConfig(pliThrottleConfig PLIThrottleConfig) ReceiverOpts { // WithAudioConfig sets up parameters for active speaker detection func WithAudioConfig(audioConfig AudioConfig) ReceiverOpts { return func(w *WebRTCReceiver) *WebRTCReceiver { - w.audioConfig = audioConfig + w.ReceiverBase.SetAudioConfig(audioConfig) return w } } func WithEnableRTPStreamRestartDetection(enable bool) ReceiverOpts { return func(w *WebRTCReceiver) *WebRTCReceiver { - w.enableRTPStreamRestartDetection = enable + w.ReceiverBase.SetEnableRTPStreamRestartDetection(enable) return w } } @@ -244,14 +81,14 @@ func WithEnableRTPStreamRestartDetection(enable bool) ReceiverOpts { // Set to 0 (disabled) by default. func WithLoadBalanceThreshold(downTracks int) ReceiverOpts { return func(w *WebRTCReceiver) *WebRTCReceiver { - w.lbThreshold = downTracks + w.ReceiverBase.SetLBThreshold(downTracks) return w } } func WithForwardStats(forwardStats *ForwardStats) ReceiverOpts { return func(w *WebRTCReceiver) *WebRTCReceiver { - w.forwardStats = forwardStats + w.ReceiverBase.SetForwardStats(forwardStats) return w } } @@ -267,111 +104,58 @@ func NewWebRTCReceiver( opts ...ReceiverOpts, ) *WebRTCReceiver { w := &WebRTCReceiver{ - logger: logger, - receiver: receiver, - trackID: livekit.TrackID(track.ID()), - streamID: track.StreamID(), - codec: track.Codec(), - codecState: ReceiverCodecStateNormal, - kind: track.Kind(), - onRTCP: onRTCP, - isRED: mime.IsMimeTypeStringRED(track.Codec().MimeType), - videoLayerMode: buffer.GetVideoLayerModeForMimeType(mime.NormalizeMimeType(track.Codec().MimeType), trackInfo), + receiver: receiver, + onRTCP: onRTCP, } + w.ReceiverBase = NewReceiverBase( + ReceiverBaseParams{ + TrackID: livekit.TrackID(track.ID()), + StreamID: track.StreamID(), + Kind: track.Kind(), + Codec: track.Codec(), + HeaderExtensions: receiver.GetParameters().HeaderExtensions, + Logger: logger, + StreamTrackerManagerConfig: streamTrackerManagerConfig, + StreamTrackerManagerListener: w, + IsSelfClosing: true, + OnClosed: w.onClosed, + }, + trackInfo, + ReceiverCodecStateNormal, + ) + for _, opt := range opts { w = opt(w) } - w.trackInfo.Store(utils.CloneProto(trackInfo)) - - w.downTrackSpreader = sfuutils.NewDownTrackSpreader[TrackSender](sfuutils.DownTrackSpreaderParams{ - Threshold: w.lbThreshold, - Logger: logger, - }) w.connectionStats = connectionquality.NewConnectionStats(connectionquality.ConnectionStatsParams{ ReceiverProvider: w, - Logger: w.logger.WithValues("direction", "up"), + Logger: logger.WithValues("direction", "up"), }) w.connectionStats.OnStatsUpdate(func(_cs *connectionquality.ConnectionStats, stat *livekit.AnalyticsStat) { if w.onStatsUpdate != nil { w.onStatsUpdate(w, stat) } }) + codec := track.Codec() w.connectionStats.Start( - mime.NormalizeMimeType(w.codec.MimeType), + mime.NormalizeMimeType(codec.MimeType), // TODO: technically not correct to declare FEC on when RED. Need the primary codec's fmtp line to check. - mime.IsMimeTypeStringRED(w.codec.MimeType) || strings.Contains(strings.ToLower(w.codec.SDPFmtpLine), "useinbandfec=1"), + mime.IsMimeTypeStringRED(codec.MimeType) || strings.Contains(strings.ToLower(codec.SDPFmtpLine), "useinbandfec=1"), ) - w.streamTrackerManager = NewStreamTrackerManager(logger, trackInfo, w.Mime(), w.codec.ClockRate, streamTrackerManagerConfig) - w.streamTrackerManager.SetListener(w) - return w } -func (w *WebRTCReceiver) TrackInfo() *livekit.TrackInfo { - return w.trackInfo.Load() -} - -func (w *WebRTCReceiver) UpdateTrackInfo(ti *livekit.TrackInfo) { - w.trackInfo.Store(utils.CloneProto(ti)) - w.streamTrackerManager.UpdateTrackInfo(ti) -} - func (w *WebRTCReceiver) OnStatsUpdate(fn func(w *WebRTCReceiver, stat *livekit.AnalyticsStat)) { w.onStatsUpdate = fn } -func (w *WebRTCReceiver) OnMaxLayerChange(fn func(mimeType mime.MimeType, maxLayer int32)) { - w.bufferMu.Lock() - w.onMaxLayerChange = fn - w.bufferMu.Unlock() -} - -func (w *WebRTCReceiver) getOnMaxLayerChange() func(mimeType mime.MimeType, maxLayer int32) { - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - - return w.onMaxLayerChange -} - func (w *WebRTCReceiver) GetConnectionScoreAndQuality() (float32, livekit.ConnectionQuality) { return w.connectionStats.GetScoreAndQuality() } -func (w *WebRTCReceiver) IsClosed() bool { - return w.closed.Load() -} - -func (w *WebRTCReceiver) SetRTT(rtt uint32) { - w.bufferMu.Lock() - if w.rtt == rtt { - w.bufferMu.Unlock() - return - } - - w.rtt = rtt - buffers := w.buffers - w.bufferMu.Unlock() - - for _, buff := range buffers { - if buff == nil { - continue - } - - buff.SetRTT(rtt) - } -} - -func (w *WebRTCReceiver) StreamID() string { - return w.streamID -} - -func (w *WebRTCReceiver) TrackID() livekit.TrackID { - return w.trackID -} - func (w *WebRTCReceiver) ssrc(layer int) uint32 { if track := w.upTracks[layer]; track != nil { return uint32(track.SSRC()) @@ -379,152 +163,44 @@ func (w *WebRTCReceiver) ssrc(layer int) uint32 { return 0 } -func (w *WebRTCReceiver) Codec() webrtc.RTPCodecParameters { - return w.codec -} - -func (w *WebRTCReceiver) Mime() mime.MimeType { - return mime.NormalizeMimeType(w.codec.MimeType) -} - -func (w *WebRTCReceiver) VideoLayerMode() livekit.VideoLayer_Mode { - return w.videoLayerMode -} - -func (w *WebRTCReceiver) HeaderExtensions() []webrtc.RTPHeaderExtensionParameter { - return w.receiver.GetParameters().HeaderExtensions -} - -func (w *WebRTCReceiver) Kind() webrtc.RTPCodecType { - return w.kind -} - func (w *WebRTCReceiver) AddUpTrack(track TrackRemote, buff *buffer.Buffer) error { - if w.closed.Load() { + if w.isClosed.Load() { return ErrReceiverClosed } layer := int32(0) if w.Kind() == webrtc.RTPCodecTypeVideo && w.videoLayerMode != livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { - layer = buffer.GetSpatialLayerForRid(w.Mime(), track.RID(), w.trackInfo.Load()) + layer = buffer.GetSpatialLayerForRid(w.Mime(), track.RID(), w.ReceiverBase.TrackInfo()) } if layer < 0 { - w.logger.Warnw( + w.ReceiverBase.Logger().Warnw( "invalid layer", nil, "rid", track.RID(), - "trackInfo", logger.Proto(w.trackInfo.Load()), + "trackInfo", logger.Proto(w.ReceiverBase.TrackInfo()), ) return ErrInvalidLayer } - buff.SetLogger(w.logger.WithValues("layer", layer)) - buff.SetAudioLevelParams(audio.AudioLevelParams{ - Config: w.audioConfig.AudioLevelConfig, - }) - buff.SetAudioLossProxying(w.audioConfig.EnableLossProxying) - buff.SetStreamRestartDetection(w.enableRTPStreamRestartDetection) - buff.OnRtcpFeedback(w.sendRTCP) - buff.OnRtcpSenderReport(func() { - srData := buff.GetSenderReportData() - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - _ = dt.HandleRTCPSenderReportData(w.codec.PayloadType, layer, srData) - }) - if rt := w.redTransformer.Load(); rt != nil { - rt.(REDTransformer).ForwardRTCPSenderReport(w.codec.PayloadType, layer, srData) - } - }) - buff.OnVideoSizeChanged(func(videoSize []buffer.VideoSize) { - w.videoSizeMu.Lock() - if w.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { - copy(w.videoSizes[:], videoSize) - } else { - w.videoSizes[layer] = videoSize[0] - } - w.logger.Debugw("video size changed", "size", w.videoSizes) - cb := w.onVideoSizeChanged - w.videoSizeMu.Unlock() - - if cb != nil { - cb() - } - }) - if w.Kind() == webrtc.RTPCodecTypeVideo && layer == 0 { - buff.OnCodecChange(w.handleCodecChange) - } - - var duration time.Duration - switch layer { - case 2: - duration = w.pliThrottleConfig.HighQuality - case 1: - duration = w.pliThrottleConfig.MidQuality - case 0: - duration = w.pliThrottleConfig.LowQuality - default: - duration = w.pliThrottleConfig.MidQuality - } - if duration != 0 { - buff.SetPLIThrottle(duration.Nanoseconds()) - } - - w.bufferMu.Lock() + w.upTracksMu.Lock() if w.upTracks[layer] != nil { - w.bufferMu.Unlock() + w.upTracksMu.Unlock() return ErrDuplicateLayer } w.upTracks[layer] = track - w.buffers[layer] = buff - rtt := w.rtt - w.bufferMu.Unlock() + w.upTracksMu.Unlock() - buff.SetRTT(rtt) - buff.SetPaused(w.streamTrackerManager.IsPaused()) - - go w.forwardRTP(layer, buff) - w.logger.Debugw("starting forwarder", "layer", layer) + w.ReceiverBase.AddBuffer(buff, layer) + buff.OnRtcpFeedback(w.sendRTCP) + w.ReceiverBase.StartBuffer(buff, layer) return nil } -// SetUpTrackPaused indicates upstream will not be sending any data. -// this will reflect the "muted" status and will pause streamtracker to ensure we don't turn off -// the layer func (w *WebRTCReceiver) SetUpTrackPaused(paused bool) { - w.streamTrackerManager.SetPaused(paused) - - w.bufferMu.RLock() - for _, buff := range w.buffers { - if buff == nil { - continue - } - - buff.SetPaused(paused) - } - w.bufferMu.RUnlock() + w.ReceiverBase.SetUpTrackPaused(paused) w.connectionStats.UpdateMute(paused) } -func (w *WebRTCReceiver) AddDownTrack(track TrackSender) error { - if w.closed.Load() { - return ErrReceiverClosed - } - - if w.downTrackSpreader.HasDownTrack(track.SubscriberID()) { - w.logger.Infow("subscriberID already exists, replacing downtrack", "subscriberID", track.SubscriberID()) - } - - track.UpTrackMaxPublishedLayerChange(w.streamTrackerManager.GetMaxPublishedLayer()) - track.UpTrackMaxTemporalLayerSeenChange(w.streamTrackerManager.GetMaxTemporalLayerSeen()) - - w.downTrackSpreader.Store(track) - w.logger.Debugw("downtrack added", "subscriberID", track.SubscriberID()) - return nil -} - -func (w *WebRTCReceiver) GetDownTracks() []TrackSender { - return w.downTrackSpreader.GetDownTracks() -} - func (w *WebRTCReceiver) notifyMaxExpectedLayer(layer int32) { ti := w.TrackInfo() if ti == nil { @@ -547,70 +223,45 @@ func (w *WebRTCReceiver) notifyMaxExpectedLayer(layer int32) { } func (w *WebRTCReceiver) SetMaxExpectedSpatialLayer(layer int32) { - w.streamTrackerManager.SetMaxExpectedSpatialLayer(layer) + w.ReceiverBase.SetMaxExpectedSpatialLayer(layer) + w.notifyMaxExpectedLayer(layer) if layer == buffer.InvalidLayerSpatial { w.connectionStats.UpdateLayerMute(true) } else { w.connectionStats.UpdateLayerMute(false) - w.connectionStats.AddLayerTransition(w.streamTrackerManager.DistanceToDesired()) + w.connectionStats.AddLayerTransition(w.ReceiverBase.StreamTrackerManager().DistanceToDesired()) } } // StreamTrackerManagerListener.OnAvailableLayersChanged func (w *WebRTCReceiver) OnAvailableLayersChanged() { - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.UpTrackLayersChange() - }) - - w.connectionStats.AddLayerTransition(w.streamTrackerManager.DistanceToDesired()) + w.connectionStats.AddLayerTransition(w.ReceiverBase.StreamTrackerManager().DistanceToDesired()) } // StreamTrackerManagerListener.OnBitrateAvailabilityChanged func (w *WebRTCReceiver) OnBitrateAvailabilityChanged() { - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.UpTrackBitrateAvailabilityChange() - }) } // StreamTrackerManagerListener.OnMaxPublishedLayerChanged func (w *WebRTCReceiver) OnMaxPublishedLayerChanged(maxPublishedLayer int32) { - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.UpTrackMaxPublishedLayerChange(maxPublishedLayer) - }) - w.notifyMaxExpectedLayer(maxPublishedLayer) - w.connectionStats.AddLayerTransition(w.streamTrackerManager.DistanceToDesired()) + w.connectionStats.AddLayerTransition(w.ReceiverBase.StreamTrackerManager().DistanceToDesired()) } // StreamTrackerManagerListener.OnMaxTemporalLayerSeenChanged func (w *WebRTCReceiver) OnMaxTemporalLayerSeenChanged(maxTemporalLayerSeen int32) { - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.UpTrackMaxTemporalLayerSeenChange(maxTemporalLayerSeen) - }) - - w.connectionStats.AddLayerTransition(w.streamTrackerManager.DistanceToDesired()) + w.connectionStats.AddLayerTransition(w.ReceiverBase.StreamTrackerManager().DistanceToDesired()) } // StreamTrackerManagerListener.OnMaxAvailableLayerChanged func (w *WebRTCReceiver) OnMaxAvailableLayerChanged(maxAvailableLayer int32) { - if onMaxLayerChange := w.getOnMaxLayerChange(); onMaxLayerChange != nil { - onMaxLayerChange(w.Mime(), maxAvailableLayer) - } } // StreamTrackerManagerListener.OnBitrateReport func (w *WebRTCReceiver) OnBitrateReport(availableLayers []int32, bitrates Bitrates) { - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.UpTrackBitrateReport(availableLayers, bitrates) - }) - - w.connectionStats.AddLayerTransition(w.streamTrackerManager.DistanceToDesired()) -} - -func (w *WebRTCReceiver) GetLayeredBitrate() ([]int32, Bitrates) { - return w.streamTrackerManager.GetLayeredBitrate() + w.connectionStats.AddLayerTransition(w.ReceiverBase.StreamTrackerManager().DistanceToDesired()) } // OnCloseHandler method to be called on remote track removed @@ -618,18 +269,8 @@ func (w *WebRTCReceiver) OnCloseHandler(fn func()) { w.onCloseHandler = fn } -// DeleteDownTrack removes a DownTrack from a Receiver -func (w *WebRTCReceiver) DeleteDownTrack(subscriberID livekit.ParticipantID) { - if w.closed.Load() { - return - } - - w.downTrackSpreader.Free(subscriberID) - w.logger.Debugw("downtrack deleted", "subscriberID", subscriberID) -} - func (w *WebRTCReceiver) sendRTCP(packets []rtcp.Packet) { - if packets == nil || w.closed.Load() { + if packets == nil || w.isClosed.Load() { return } @@ -638,93 +279,10 @@ func (w *WebRTCReceiver) sendRTCP(packets []rtcp.Packet) { } } -func (w *WebRTCReceiver) SendPLI(layer int32, force bool) { - // SVC-TODO : should send LRR (Layer Refresh Request) instead of PLI - buff := w.getBuffer(layer) - if buff == nil { - return - } - - buff.SendPLI(force) -} - -func (w *WebRTCReceiver) getBuffer(layer int32) *buffer.Buffer { - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - - return w.getBufferLocked(layer) -} - -func (w *WebRTCReceiver) getBufferLocked(layer int32) *buffer.Buffer { - // for svc codecs, use layer = 0 always. - // spatial layers are in-built and handled by single buffer - if w.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { - layer = 0 - } - - if layer < 0 || int(layer) >= len(w.buffers) { - return nil - } - - return w.buffers[layer] -} - -func (w *WebRTCReceiver) ReadRTP(buf []byte, layer uint8, esn uint64) (int, error) { - b := w.getBuffer(int32(layer)) - if b == nil { - return 0, ErrBufferNotFound - } - - return b.GetPacket(buf, esn) -} - -func (w *WebRTCReceiver) GetTrackStats() *livekit.RTPStats { - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - - stats := make([]*livekit.RTPStats, 0, len(w.buffers)) - for _, buff := range w.buffers { - if buff == nil { - continue - } - - sswl := buff.GetStats() - if sswl == nil { - continue - } - - stats = append(stats, sswl) - } - - return rtpstats.AggregateRTPStats(stats) -} - -func (w *WebRTCReceiver) GetAudioLevel() (float64, bool) { - if w.Kind() == webrtc.RTPCodecTypeVideo { - return 0, false - } - - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - - for _, buff := range w.buffers { - if buff == nil { - continue - } - - return buff.GetAudioLevel() - } - - return 0, false -} - func (w *WebRTCReceiver) GetDeltaStats() map[uint32]*buffer.StreamStatsWithLayers { - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - - deltaStats := make(map[uint32]*buffer.StreamStatsWithLayers, len(w.buffers)) - - for layer, buff := range w.buffers { + buffers := w.ReceiverBase.GetAllBuffers() + deltaStats := make(map[uint32]*buffer.StreamStatsWithLayers, len(buffers)) + for layer, buff := range buffers { if buff == nil { continue } @@ -746,11 +304,9 @@ func (w *WebRTCReceiver) GetDeltaStats() map[uint32]*buffer.StreamStatsWithLayer } func (w *WebRTCReceiver) GetLastSenderReportTime() time.Time { - w.bufferMu.RLock() - defer w.bufferMu.RUnlock() - + buffers := w.ReceiverBase.GetAllBuffers() latestSRTime := time.Time{} - for _, buff := range w.buffers { + for _, buff := range buffers { if buff == nil { continue } @@ -764,161 +320,22 @@ func (w *WebRTCReceiver) GetLastSenderReportTime() time.Time { return latestSRTime } -func (w *WebRTCReceiver) forwardRTP(layer int32, buff *buffer.Buffer) { - numPacketsForwarded := 0 - numPacketsDropped := 0 - defer func() { - w.closeOnce.Do(func() { - w.closed.Store(true) - w.closeTracks() - if rt := w.redTransformer.Load(); rt != nil { - rt.(REDTransformer).Close() - } - }) - - w.streamTrackerManager.RemoveTracker(layer) - if w.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { - w.streamTrackerManager.RemoveAllTrackers() - } - - w.logger.Debugw( - "closing forwarder", - "layer", layer, - "numPacketsForwarded", numPacketsForwarded, - "numPacketsDropped", numPacketsDropped, - ) - }() - - var spatialTrackers [buffer.DefaultMaxLayerSpatial + 1]streamtracker.StreamTrackerWorker - if layer < 0 || int(layer) >= len(spatialTrackers) { - w.logger.Errorw("invalid layer", nil, "layer", layer) - return - } - - pktBuf := make([]byte, bucket.RTPMaxPktSize) - w.logger.Debugw("starting forwarding", "layer", layer) - for { - pkt, err := buff.ReadExtended(pktBuf) - if err == io.EOF { - return - } - dequeuedAt := mono.UnixNano() - - if pkt.IsRestart { - w.logger.Infow("stream restarted", "layer", layer) - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - dt.ReceiverRestart() - }) - - if rt := w.redTransformer.Load(); rt != nil { - rt.(REDTransformer).OnStreamRestart() - } - } - - if pkt.Packet.PayloadType != uint8(w.codec.PayloadType) { - // drop packets as we don't support codec fallback directly - w.logger.Debugw( - "dropping packet - payload mismatch", - "packetPayloadType", pkt.Packet.PayloadType, - "payloadType", w.codec.PayloadType, - ) - numPacketsDropped++ - continue - } - - spatialLayer := layer - if pkt.Spatial >= 0 { - // svc packet, take spatial layer info from packet - spatialLayer = pkt.Spatial - } - if int(spatialLayer) >= len(spatialTrackers) { - w.logger.Errorw( - "unexpected spatial layer", nil, - "spatialLayer", spatialLayer, - "pktSpatialLayer", pkt.Spatial, - ) - numPacketsDropped++ - continue - } - - var writeCount atomic.Int32 - w.downTrackSpreader.Broadcast(func(dt TrackSender) { - writeCount.Add(dt.WriteRTP(pkt, spatialLayer)) - }) - - if rt := w.redTransformer.Load(); rt != nil { - writeCount.Add(rt.(REDTransformer).ForwardRTP(pkt, spatialLayer)) - } - - // track delay/jitter - if writeCount.Load() > 0 && w.forwardStats != nil && !pkt.IsBuffered { - if latency, isHigh := w.forwardStats.Update(pkt.Arrival, mono.UnixNano()); isHigh { - w.logger.Debugw( - "high forwarding latency", - "latency", time.Duration(latency), - "queuingLatency", time.Duration(dequeuedAt-pkt.Arrival), - "writeCount", writeCount.Load(), - "isOutOfOrder", pkt.IsOutOfOrder, - "layer", layer, - ) - } - } - - // track video layers - if w.Kind() == webrtc.RTPCodecTypeVideo { - if spatialTrackers[spatialLayer] == nil { - spatialTrackers[spatialLayer] = w.streamTrackerManager.GetTracker(spatialLayer) - if spatialTrackers[spatialLayer] == nil { - if w.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM && pkt.DependencyDescriptor != nil { - w.streamTrackerManager.AddDependencyDescriptorTrackers() - } - spatialTrackers[spatialLayer] = w.streamTrackerManager.AddTracker(spatialLayer) - } - } - if spatialTrackers[spatialLayer] != nil { - spatialTrackers[spatialLayer].Observe( - pkt.Temporal, - len(pkt.RawPacket), - len(pkt.Packet.Payload), - pkt.Packet.Marker, - pkt.Packet.Timestamp, - pkt.DependencyDescriptor, - ) - } - } - - numPacketsForwarded++ - - buffer.ReleaseExtPacket(pkt) - } -} - -// closeTracks close all tracks from Receiver -func (w *WebRTCReceiver) closeTracks() { +func (w *WebRTCReceiver) onClosed() { w.connectionStats.Close() - w.streamTrackerManager.Close() - - closeTrackSenders(w.downTrackSpreader.ResetAndGetDownTracks()) if w.onCloseHandler != nil { w.onCloseHandler() } } -func (w *WebRTCReceiver) DebugInfo() map[string]interface{} { - var videoLayerMode livekit.VideoLayer_Mode - if ti := w.trackInfo.Load(); ti != nil { - videoLayerMode = buffer.GetVideoLayerModeForMimeType(w.Mime(), ti) - } - info := map[string]interface{}{ - "VideoLayerMode": videoLayerMode.String(), - } +func (w *WebRTCReceiver) DebugInfo() map[string]any { + info := w.ReceiverBase.DebugInfo() - w.bufferMu.RLock() - upTrackInfo := make([]map[string]interface{}, 0, len(w.upTracks)) + w.upTracksMu.Lock() + upTrackInfo := make([]map[string]any, 0, len(w.upTracks)) for layer, ut := range w.upTracks { if ut != nil { - upTrackInfo = append(upTrackInfo, map[string]interface{}{ + upTrackInfo = append(upTrackInfo, map[string]any{ "Layer": layer, "SSRC": ut.SSRC(), "Msid": ut.Msid(), @@ -926,144 +343,10 @@ func (w *WebRTCReceiver) DebugInfo() map[string]interface{} { }) } } - w.bufferMu.RUnlock() + w.upTracksMu.Unlock() info["UpTracks"] = upTrackInfo return info } -func (w *WebRTCReceiver) GetPrimaryReceiverForRed() TrackReceiver { - w.bufferMu.Lock() - defer w.bufferMu.Unlock() - - if !w.isRED || w.closed.Load() { - return w - } - - rt := w.redTransformer.Load() - if rt == nil { - pr := NewRedPrimaryReceiver(w, sfuutils.DownTrackSpreaderParams{ - Threshold: w.lbThreshold, - Logger: w.logger, - }) - w.redTransformer.Store(pr) - return pr - } else { - if pr, ok := rt.(*RedPrimaryReceiver); ok { - return pr - } - } - return nil -} - -func (w *WebRTCReceiver) GetRedReceiver() TrackReceiver { - w.bufferMu.Lock() - defer w.bufferMu.Unlock() - - if w.isRED || w.closed.Load() { - return w - } - - rt := w.redTransformer.Load() - if rt == nil { - pr := NewRedReceiver(w, sfuutils.DownTrackSpreaderParams{ - Threshold: w.lbThreshold, - Logger: w.logger, - }) - w.redTransformer.Store(pr) - return pr - } else { - if pr, ok := rt.(*RedReceiver); ok { - return pr - } - } - return nil -} - -func (w *WebRTCReceiver) GetTemporalLayerFpsForSpatial(layer int32) []float32 { - b := w.getBuffer(layer) - if b == nil { - return nil - } - - if w.videoLayerMode != livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { - return b.GetTemporalLayerFpsForSpatial(0) - } - - return b.GetTemporalLayerFpsForSpatial(layer) -} - -func (w *WebRTCReceiver) AddOnReady(fn func()) { - // webRTCReceiver is always ready after created - fn() -} - -func (w *WebRTCReceiver) handleCodecChange(newCodec webrtc.RTPCodecParameters) { - // we don't support the codec fallback directly, set the codec state to invalid once it happens - w.SetCodecState(ReceiverCodecStateInvalid) -} - -func (w *WebRTCReceiver) AddOnCodecStateChange(f func(webrtc.RTPCodecParameters, ReceiverCodecState)) { - w.codecStateLock.Lock() - w.onCodecStateChange = append(w.onCodecStateChange, f) - w.codecStateLock.Unlock() -} - -func (w *WebRTCReceiver) CodecState() ReceiverCodecState { - w.codecStateLock.Lock() - defer w.codecStateLock.Unlock() - - return w.codecState -} - -func (w *WebRTCReceiver) SetCodecState(state ReceiverCodecState) { - w.codecStateLock.Lock() - if w.codecState == state || w.codecState == ReceiverCodecStateInvalid { - w.codecStateLock.Unlock() - return - } - - w.codecState = state - fns := w.onCodecStateChange - w.codecStateLock.Unlock() - - for _, f := range fns { - f(w.codec, state) - } -} - -func (w *WebRTCReceiver) VideoSizes() []buffer.VideoSize { - var sizes []buffer.VideoSize - w.videoSizeMu.RLock() - defer w.videoSizeMu.RUnlock() - for _, v := range w.videoSizes { - if v.Width == 0 || v.Height == 0 { - break - } - sizes = append(sizes, v) - } - - return sizes -} - -func (w *WebRTCReceiver) OnVideoSizeChanged(f func()) { - w.videoSizeMu.Lock() - w.onVideoSizeChanged = f - w.videoSizeMu.Unlock() -} - // ----------------------------------------------------------- - -// closes all track senders in parallel, returns when all are closed -func closeTrackSenders(senders []TrackSender) { - wg := sync.WaitGroup{} - for _, dt := range senders { - dt := dt - wg.Add(1) - go func() { - defer wg.Done() - dt.Close() - }() - } - wg.Wait() -} diff --git a/pkg/sfu/receiver_base.go b/pkg/sfu/receiver_base.go new file mode 100644 index 000000000..355c6a00a --- /dev/null +++ b/pkg/sfu/receiver_base.go @@ -0,0 +1,1111 @@ +// 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 sfu + +import ( + "errors" + "fmt" + "io" + "slices" + "strings" + "sync" + "time" + + "github.com/pion/webrtc/v4" + "go.uber.org/atomic" + + "github.com/livekit/mediatransportutil/pkg/bucket" + "github.com/livekit/protocol/livekit" + "github.com/livekit/protocol/logger" + "github.com/livekit/protocol/utils" + "github.com/livekit/protocol/utils/mono" + + "github.com/livekit/livekit-server/pkg/sfu/audio" + "github.com/livekit/livekit-server/pkg/sfu/buffer" + "github.com/livekit/livekit-server/pkg/sfu/mime" + "github.com/livekit/livekit-server/pkg/sfu/rtpstats" + "github.com/livekit/livekit-server/pkg/sfu/streamtracker" + sfuutils "github.com/livekit/livekit-server/pkg/sfu/utils" +) + +var ( + ErrReceiverClosed = errors.New("receiver closed") + ErrDownTrackAlreadyExist = errors.New("DownTrack already exist") + ErrDuplicateLayer = errors.New("duplicate layer") + ErrInvalidLayer = errors.New("invalid layer") +) + +// -------------------------------------- + +type PLIThrottleConfig struct { + LowQuality time.Duration `yaml:"low_quality,omitempty"` + MidQuality time.Duration `yaml:"mid_quality,omitempty"` + HighQuality time.Duration `yaml:"high_quality,omitempty"` +} + +var ( + DefaultPLIThrottleConfig = PLIThrottleConfig{ + LowQuality: 500 * time.Millisecond, + MidQuality: time.Second, + HighQuality: time.Second, + } +) + +// -------------------------------------- + +type AudioConfig struct { + audio.AudioLevelConfig `yaml:",inline"` + + // enable red encoding downtrack for opus only audio up track + ActiveREDEncoding bool `yaml:"active_red_encoding,omitempty"` + // enable proxying weakest subscriber loss to publisher in RTCP Receiver Report + EnableLossProxying bool `yaml:"enable_loss_proxying,omitempty"` +} + +var ( + DefaultAudioConfig = AudioConfig{ + AudioLevelConfig: audio.DefaultAudioLevelConfig, + } +) + +// -------------------------------------- + +type AudioLevelHandle func(level uint8, duration uint32) + +// -------------------------------------- + +type Bitrates [buffer.DefaultMaxLayerSpatial + 1][buffer.DefaultMaxLayerTemporal + 1]int64 + +// -------------------------------------- + +type ReceiverCodecState int + +const ( + ReceiverCodecStateNormal ReceiverCodecState = iota + ReceiverCodecStateSuspended + ReceiverCodecStateInvalid +) + +// -------------------------------------- + +// TrackReceiver defines an interface receive media from remote peer +type TrackReceiver interface { + TrackID() livekit.TrackID + StreamID() string + + // returns the initial codec of the receiver, it is determined by the track's codec + // and will not change if the codec changes during the session (publisher changes codec) + Codec() webrtc.RTPCodecParameters + Mime() mime.MimeType + VideoLayerMode() livekit.VideoLayer_Mode + HeaderExtensions() []webrtc.RTPHeaderExtensionParameter + IsClosed() bool + + ReadRTP(buf []byte, layer uint8, esn uint64) (int, error) + GetLayeredBitrate() ([]int32, Bitrates) + + GetAudioLevel() (float64, bool) + + SendPLI(layer int32, force bool) + + SetUpTrackPaused(paused bool) + SetMaxExpectedSpatialLayer(layer int32) + + AddDownTrack(track TrackSender) error + DeleteDownTrack(participantID livekit.ParticipantID) + GetDownTracks() []TrackSender + + DebugInfo() map[string]any + + TrackInfo() *livekit.TrackInfo + UpdateTrackInfo(ti *livekit.TrackInfo) + + // Get primary receiver if this receiver represents a RED codec; otherwise it will return itself + GetPrimaryReceiverForRed() TrackReceiver + + // Get red receiver for primary codec, used by forward red encodings for opus only codec + GetRedReceiver() TrackReceiver + + GetTemporalLayerFpsForSpatial(layer int32) []float32 + + GetTrackStats() *livekit.RTPStats + + // AddOnReady adds a function to be called when the receiver is ready, the callback + // could be called immediately if the receiver is ready when the callback is added + AddOnReady(func()) + + AddOnCodecStateChange(func(webrtc.RTPCodecParameters, ReceiverCodecState)) + CodecState() ReceiverCodecState + + // VideoSizes returns the video size parsed from rtp packet for each spatial layer. + VideoSizes() []buffer.VideoSize +} + +// -------------------------------------- + +type REDTransformer interface { + ForwardRTP(pkt *buffer.ExtPacket, spatialLayer int32) int32 + ForwardRTCPSenderReport( + payloadType webrtc.PayloadType, + layer int32, + publisherSRData *livekit.RTCPSenderReportState, + ) + GetDownTracks() []TrackSender + HasDownTracks() bool + ResyncDownTracks() + OnStreamRestart() + CanClose() bool + Close() +} + +// -------------------------------------- + +type ReceiverBaseParams struct { + TrackID livekit.TrackID + StreamID string + Kind webrtc.RTPCodecType + Codec webrtc.RTPCodecParameters + HeaderExtensions []webrtc.RTPHeaderExtensionParameter + Logger logger.Logger + StreamTrackerManagerConfig StreamTrackerManagerConfig + StreamTrackerManagerListener StreamTrackerManagerListener + IsSelfClosing bool + OnClosed func() +} + +type ReceiverBase struct { + params ReceiverBaseParams + + pliThrottleConfig PLIThrottleConfig + audioConfig AudioConfig + enableRTPStreamRestartDetection bool + lbThreshold int + forwardStats *ForwardStats + + codecStateLock sync.Mutex + codecState ReceiverCodecState + onCodecStateChange []func(webrtc.RTPCodecParameters, ReceiverCodecState) + + isRED bool + videoLayerMode livekit.VideoLayer_Mode + + bufferMu sync.RWMutex + buffers [buffer.DefaultMaxLayerSpatial + 1]buffer.BufferProvider + trackInfo *livekit.TrackInfo + + videoSizeMu sync.RWMutex + videoSizes [buffer.DefaultMaxLayerSpatial + 1]buffer.VideoSize + onVideoSizeChanged func() + + rtt uint32 + + streamTrackerManager *StreamTrackerManager + + downTrackSpreader *sfuutils.DownTrackSpreader[TrackSender] + + onMaxLayerChange func(mimeType mime.MimeType, maxLayer int32) + + redTransformer atomic.Value // redTransformer interface + + isClosed atomic.Bool +} + +func NewReceiverBase(params ReceiverBaseParams, trackInfo *livekit.TrackInfo, codecState ReceiverCodecState) *ReceiverBase { + r := &ReceiverBase{ + params: params, + codecState: codecState, + isRED: mime.IsMimeTypeStringRED(params.Codec.MimeType), + trackInfo: utils.CloneProto(trackInfo), + videoLayerMode: buffer.GetVideoLayerModeForMimeType(mime.NormalizeMimeType(params.Codec.MimeType), trackInfo), + } + + r.downTrackSpreader = sfuutils.NewDownTrackSpreader[TrackSender](sfuutils.DownTrackSpreaderParams{ + Threshold: r.lbThreshold, + Logger: params.Logger, + }) + + r.streamTrackerManager = NewStreamTrackerManager( + params.Logger, + trackInfo, + r.Mime(), + r.params.Codec.ClockRate, + params.StreamTrackerManagerConfig, + ) + r.streamTrackerManager.SetListener(r) + + return r +} + +func (r *ReceiverBase) Close(reason string, clearBuffers bool) { + if r.isClosed.Swap(true) { + return + } + + if clearBuffers { + r.ClearAllBuffers(reason) + } + r.streamTrackerManager.Close() + + closeTrackSenders(r.downTrackSpreader.ResetAndGetDownTracks()) + + if rt := r.redTransformer.Load(); rt != nil { + rt.(REDTransformer).Close() + } + + if r.params.OnClosed != nil { + r.params.OnClosed() + } +} + +func (r *ReceiverBase) CanClose() bool { + if r.IsClosed() { + return true + } + + if r.downTrackSpreader.DownTrackCount() != 0 { + return false + } + + if rt := r.redTransformer.Load(); rt != nil { + return rt.(REDTransformer).CanClose() + } + + return true +} + +func (r *ReceiverBase) SetPLIThrottleConfig(pliThrottleConfig PLIThrottleConfig) { + r.pliThrottleConfig = pliThrottleConfig +} + +func (r *ReceiverBase) SetAudioConfig(audioConfig AudioConfig) { + r.audioConfig = audioConfig +} + +func (r *ReceiverBase) SetEnableRTPStreamRestartDetection(enableRTPStremRestartDetection bool) { + r.enableRTPStreamRestartDetection = enableRTPStremRestartDetection +} + +func (r *ReceiverBase) SetLBThreshold(lbThreshold int) { + r.lbThreshold = lbThreshold +} + +func (r *ReceiverBase) SetForwardStats(forwardStats *ForwardStats) { + r.forwardStats = forwardStats +} + +func (r *ReceiverBase) Logger() logger.Logger { + return r.params.Logger +} + +func (r *ReceiverBase) TrackInfo() *livekit.TrackInfo { + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + return utils.CloneProto(r.trackInfo) +} + +func (r *ReceiverBase) UpdateTrackInfo(ti *livekit.TrackInfo) { + r.bufferMu.Lock() + existingVersion := utils.TimedVersionFromProto(r.trackInfo.Version) + updateVersion := utils.TimedVersionFromProto(ti.Version) + if updateVersion.Compare(existingVersion) < 0 { + r.bufferMu.Unlock() + r.params.Logger.Debugw( + "not updating to older version", + "existing", logger.Proto(r.trackInfo), + "updated", logger.Proto(ti), + ) + return + } + + shouldResync := utils.TimedVersionFromProto(r.trackInfo.Version) != utils.TimedVersionFromProto(ti.Version) + if shouldResync { + r.params.Logger.Debugw( + "updating track info", + "existing", logger.Proto(r.trackInfo), + "updated", logger.Proto(ti), + "shouldResync", shouldResync, + ) + } + r.trackInfo = utils.CloneProto(ti) + // MUTABLE-TRACKINFO-TODO: notify buffers, buffers may need to resize retransmission buffer if there is layer change + + if shouldResync { + r.resyncLocked("update-track-info") + } + r.bufferMu.Unlock() + + r.streamTrackerManager.UpdateTrackInfo(ti) +} + +func (r *ReceiverBase) resyncLocked(reason string) { + // resync to avoid gaps in the forwarded sequence number + r.params.Logger.Debugw("resync receiver", "reason", reason) + r.clearAllBuffersLocked("resync") + + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.Resync() + }) + if rt := r.redTransformer.Load(); rt != nil { + rt.(REDTransformer).ResyncDownTracks() + } +} + +func (r *ReceiverBase) OnMaxLayerChange(fn func(mimeType mime.MimeType, maxLayer int32)) { + r.bufferMu.Lock() + r.onMaxLayerChange = fn + r.bufferMu.Unlock() +} + +func (r *ReceiverBase) getOnMaxLayerChange() func(mimeType mime.MimeType, maxLayer int32) { + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + return r.onMaxLayerChange +} + +func (r *ReceiverBase) IsClosed() bool { + return r.isClosed.Load() +} + +func (r *ReceiverBase) SetRTT(rtt uint32) { + r.bufferMu.Lock() + if r.rtt == rtt || rtt == 0 { + r.bufferMu.Unlock() + return + } + + r.rtt = rtt + buffers := r.buffers + r.bufferMu.Unlock() + + for _, buff := range buffers { + if buff == nil { + continue + } + + buff.SetRTT(rtt) + } +} + +func (r *ReceiverBase) TrackID() livekit.TrackID { + return r.params.TrackID +} + +func (r *ReceiverBase) StreamID() string { + return r.params.StreamID +} + +func (r *ReceiverBase) Codec() webrtc.RTPCodecParameters { + return r.params.Codec +} + +func (r *ReceiverBase) Mime() mime.MimeType { + return mime.NormalizeMimeType(r.params.Codec.MimeType) +} + +func (r *ReceiverBase) VideoLayerMode() livekit.VideoLayer_Mode { + return r.videoLayerMode +} + +func (r *ReceiverBase) HeaderExtensions() []webrtc.RTPHeaderExtensionParameter { + return r.params.HeaderExtensions +} + +func (r *ReceiverBase) Kind() webrtc.RTPCodecType { + return r.params.Kind +} + +func (r *ReceiverBase) StreamTrackerManager() *StreamTrackerManager { + return r.streamTrackerManager +} + +// SetUpTrackPaused indicates upstream will not be sending any data. +// this will reflect the "muted" status and will pause streamtracker to ensure we don't turn off +// the layer +func (r *ReceiverBase) SetUpTrackPaused(paused bool) { + r.streamTrackerManager.SetPaused(paused) + + r.bufferMu.RLock() + for _, buff := range r.buffers { + if buff == nil { + continue + } + + buff.SetPaused(paused) + } + r.bufferMu.RUnlock() +} + +func (r *ReceiverBase) AddDownTrack(track TrackSender) error { + if r.IsClosed() { + return ErrReceiverClosed + } + + if r.downTrackSpreader.HasDownTrack(track.SubscriberID()) { + r.params.Logger.Infow("subscriberID already exists, replacing downtrack", "subscriberID", track.SubscriberID()) + } + + track.UpTrackMaxPublishedLayerChange(r.streamTrackerManager.GetMaxPublishedLayer()) + track.UpTrackMaxTemporalLayerSeenChange(r.streamTrackerManager.GetMaxTemporalLayerSeen()) + + r.downTrackSpreader.Store(track) + r.params.Logger.Debugw("downtrack added", "subscriberID", track.SubscriberID()) + return nil +} + +func (r *ReceiverBase) DeleteDownTrack(subscriberID livekit.ParticipantID) { + r.downTrackSpreader.Free(subscriberID) + r.params.Logger.Debugw("downtrack deleted", "subscriberID", subscriberID) +} + +func (r *ReceiverBase) GetDownTracks() []TrackSender { + downTracks := r.downTrackSpreader.GetDownTracks() + if rt := r.redTransformer.Load(); rt != nil { + downTracks = append(downTracks, rt.(REDTransformer).GetDownTracks()...) + } + return downTracks +} + +func (r *ReceiverBase) SetMaxExpectedSpatialLayer(layer int32) { + prevMax := r.streamTrackerManager.SetMaxExpectedSpatialLayer(layer) + r.params.Logger.Debugw("max expected layer change", "layer", layer, "prevMax", prevMax) + + r.bufferMu.RLock() + // stop key frame seeders of stopped layers + for idx := layer + 1; idx <= prevMax; idx++ { + if r.buffers[idx] != nil { + r.buffers[idx].StopKeyFrameSeeder() + } + } + + // start key frame seeders of newly expected layers + for idx := prevMax + 1; idx <= layer; idx++ { + if r.buffers[idx] != nil { + r.buffers[idx].StartKeyFrameSeeder() + } + } + r.bufferMu.RUnlock() +} + +// StreamTrackerManagerListener.OnAvailableLayersChanged +func (r *ReceiverBase) OnAvailableLayersChanged() { + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.UpTrackLayersChange() + }) + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnAvailableLayersChanged() + } +} + +// StreamTrackerManagerListener.OnBitrateAvailabilityChanged +func (r *ReceiverBase) OnBitrateAvailabilityChanged() { + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.UpTrackBitrateAvailabilityChange() + }) + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnBitrateAvailabilityChanged() + } +} + +// StreamTrackerManagerListener.OnMaxPublishedLayerChanged +func (r *ReceiverBase) OnMaxPublishedLayerChanged(maxPublishedLayer int32) { + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.UpTrackMaxPublishedLayerChange(maxPublishedLayer) + }) + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnMaxPublishedLayerChanged(maxPublishedLayer) + } +} + +// StreamTrackerManagerListener.OnMaxTemporalLayerSeenChanged +func (r *ReceiverBase) OnMaxTemporalLayerSeenChanged(maxTemporalLayerSeen int32) { + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.UpTrackMaxTemporalLayerSeenChange(maxTemporalLayerSeen) + }) + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnMaxTemporalLayerSeenChanged(maxTemporalLayerSeen) + } +} + +// StreamTrackerManagerListener.OnMaxAvailableLayerChanged +func (r *ReceiverBase) OnMaxAvailableLayerChanged(maxAvailableLayer int32) { + if onMaxLayerChange := r.getOnMaxLayerChange(); onMaxLayerChange != nil { + onMaxLayerChange(r.Mime(), maxAvailableLayer) + } + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnMaxAvailableLayerChanged(maxAvailableLayer) + } +} + +// StreamTrackerManagerListener.OnBitrateReport +func (r *ReceiverBase) OnBitrateReport(availableLayers []int32, bitrates Bitrates) { + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.UpTrackBitrateReport(availableLayers, bitrates) + }) + + if r.params.StreamTrackerManagerListener != nil { + r.params.StreamTrackerManagerListener.OnBitrateReport(availableLayers, bitrates) + } +} + +func (r *ReceiverBase) GetLayeredBitrate() ([]int32, Bitrates) { + return r.streamTrackerManager.GetLayeredBitrate() +} + +func (r *ReceiverBase) SendPLI(layer int32, force bool) { + // SVC-TODO : should send LRR (Layer Refresh Request) instead of PLI + buff := r.getBuffer(layer) + if buff == nil { + return + } + + buff.SendPLI(force) +} + +func (r *ReceiverBase) getBuffer(layer int32) buffer.BufferProvider { + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + return r.getBufferLocked(layer) +} + +func (r *ReceiverBase) getBufferLocked(layer int32) buffer.BufferProvider { + // for svc codecs, use layer = 0 always. + // spatial layers are in-built and handled by single buffer + if r.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { + layer = 0 + } + + if layer < 0 || int(layer) >= len(r.buffers) { + return nil + } + + return r.buffers[layer] +} + +func (r *ReceiverBase) GetOrCreateBuffer( + layer int32, + creatorFn func() (buffer.BufferProvider, error), +) (buffer.BufferProvider, bool) { + r.bufferMu.Lock() + + if r.IsClosed() { + r.bufferMu.Unlock() + return nil, false + } + + if buff := r.getBufferLocked(layer); buff != nil { + r.bufferMu.Unlock() + return buff, false + } + + buff, err := creatorFn() + if err != nil { + r.bufferMu.Unlock() + r.params.Logger.Errorw("could not create buffer", err) + return nil, false + } + + rtt := r.rtt + r.bufferMu.Unlock() + + r.setupBuffer(buff, layer, rtt) + return buff, true +} + +func (r *ReceiverBase) setupBuffer(buff buffer.BufferProvider, layer int32, rtt uint32) { + buff.SetLogger(r.params.Logger.WithValues("layer", layer)) + buff.SetAudioLevelParams(audio.AudioLevelParams{ + Config: r.audioConfig.AudioLevelConfig, + }) + buff.SetStreamRestartDetection(r.enableRTPStreamRestartDetection) + buff.OnRtcpSenderReport(func() { + srData := buff.GetSenderReportData() + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + _ = dt.HandleRTCPSenderReportData(r.params.Codec.PayloadType, layer, srData) + }) + + if rt := r.redTransformer.Load(); rt != nil { + rt.(REDTransformer).ForwardRTCPSenderReport(r.params.Codec.PayloadType, layer, srData) + } + }) + buff.OnVideoSizeChanged(func(videoSize []buffer.VideoSize) { + r.videoSizeMu.Lock() + if r.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { + copy(r.videoSizes[:], videoSize) + } else { + r.videoSizes[layer] = videoSize[0] + } + r.params.Logger.Debugw("video size changed", "size", r.videoSizes) + cb := r.onVideoSizeChanged + r.videoSizeMu.Unlock() + + if cb != nil { + cb() + } + }) + if r.Kind() == webrtc.RTPCodecTypeVideo && layer == 0 { + buff.OnCodecChange(r.handleCodecChange) + } + + var duration time.Duration + switch layer { + case 2: + duration = r.pliThrottleConfig.HighQuality + case 1: + duration = r.pliThrottleConfig.MidQuality + case 0: + duration = r.pliThrottleConfig.LowQuality + default: + duration = r.pliThrottleConfig.MidQuality + } + if duration != 0 { + buff.SetPLIThrottle(duration.Nanoseconds()) + } + + buff.SetRTT(rtt) + buff.SetPaused(r.streamTrackerManager.IsPaused()) +} + +func (r *ReceiverBase) AddBuffer(buff buffer.BufferProvider, layer int32) { + r.bufferMu.Lock() + r.buffers[layer] = buff + rtt := r.rtt + r.bufferMu.Unlock() + + r.setupBuffer(buff, layer, rtt) +} + +func (r *ReceiverBase) StartBuffer(buff buffer.BufferProvider, layer int32) { + r.params.Logger.Debugw("starting forwarder", "layer", layer) + go r.forwardRTP(layer, buff) +} + +func (r *ReceiverBase) GetAllBuffers() [buffer.DefaultMaxLayerSpatial + 1]buffer.BufferProvider { + buffers := [buffer.DefaultMaxLayerSpatial + 1]buffer.BufferProvider{} + + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + for i := range buffers { + buffers[i] = r.buffers[i] + } + return buffers +} + +func (r *ReceiverBase) ClearAllBuffers(reason string) { + r.bufferMu.Lock() + defer r.bufferMu.Unlock() + + r.clearAllBuffersLocked(reason) +} + +func (r *ReceiverBase) clearAllBuffersLocked(reason string) { + for idx := range len(r.buffers) { + if r.buffers[idx] != nil { + r.buffers[idx].CloseWithReason(reason) + } + r.buffers[idx] = nil + } + + r.streamTrackerManager.RemoveAllTrackers() +} + +func (r *ReceiverBase) ReadRTP(buf []byte, layer uint8, esn uint64) (int, error) { + b := r.getBuffer(int32(layer)) + if b == nil { + return 0, bucket.ErrPacketMismatch + } + + return b.GetPacket(buf, esn) +} + +func (r *ReceiverBase) GetTrackStats() *livekit.RTPStats { + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + allStats := make([]*livekit.RTPStats, 0, len(r.buffers)) + for _, buff := range r.buffers { + if buff == nil { + continue + } + + stats := buff.GetStats() + if stats == nil { + continue + } + + allStats = append(allStats, stats) + } + + return rtpstats.AggregateRTPStats(allStats) +} + +func (r *ReceiverBase) GetAudioLevel() (float64, bool) { + if r.Kind() == webrtc.RTPCodecTypeVideo { + return 0, false + } + + r.bufferMu.RLock() + defer r.bufferMu.RUnlock() + + for _, buff := range r.buffers { + if buff == nil { + continue + } + + return buff.GetAudioLevel() + } + + return 0, false +} + +func (r *ReceiverBase) forwardRTP(layer int32, buff buffer.BufferProvider) { + numPacketsForwarded := 0 + numPacketsDropped := 0 + defer func() { + if r.params.IsSelfClosing { + r.Close("forwarder-done", false) + + r.streamTrackerManager.RemoveTracker(layer) + if r.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { + r.streamTrackerManager.RemoveAllTrackers() + } + } + + r.params.Logger.Debugw( + "closing forwarder", + "layer", layer, + "numPacketsForwarded", numPacketsForwarded, + "numPacketsDropped", numPacketsDropped, + ) + }() + + var spatialTrackers [buffer.DefaultMaxLayerSpatial + 1]streamtracker.StreamTrackerWorker + if layer < 0 || int(layer) >= len(spatialTrackers) { + r.params.Logger.Errorw("invalid layer", nil, "layer", layer) + return + } + + pktBuf := make([]byte, bucket.RTPMaxPktSize) + r.params.Logger.Debugw("starting forwarding", "layer", layer) + for { + pkt, err := buff.ReadExtended(pktBuf) + if err == io.EOF { + return + } + dequeuedAt := mono.UnixNano() + + if pkt.IsRestart { + r.params.Logger.Infow("stream restarted", "layer", layer) + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + dt.ReceiverRestart() + }) + + if rt := r.redTransformer.Load(); rt != nil { + rt.(REDTransformer).OnStreamRestart() + } + } + + if pkt.Packet.PayloadType != uint8(r.params.Codec.PayloadType) { + // drop packets as we don't support codec fallback directly + r.params.Logger.Debugw( + "dropping packet - payload mismatch", + "packetPayloadType", pkt.Packet.PayloadType, + "payloadType", r.params.Codec.PayloadType, + ) + numPacketsDropped++ + continue + } + + spatialLayer := layer + if pkt.Spatial >= 0 { + // svc packet, take spatial layer info from packet + spatialLayer = pkt.Spatial + } + if int(spatialLayer) >= len(spatialTrackers) { + r.params.Logger.Errorw( + "unexpected spatial layer", nil, + "spatialLayer", spatialLayer, + "pktSpatialLayer", pkt.Spatial, + ) + numPacketsDropped++ + continue + } + + var writeCount atomic.Int32 + r.downTrackSpreader.Broadcast(func(dt TrackSender) { + writeCount.Add(dt.WriteRTP(pkt, spatialLayer)) + }) + + if rt := r.redTransformer.Load(); rt != nil { + writeCount.Add(rt.(REDTransformer).ForwardRTP(pkt, spatialLayer)) + } + + // track delay/jitter + if writeCount.Load() > 0 && r.forwardStats != nil && !pkt.IsBuffered { + if latency, isHigh := r.forwardStats.Update(pkt.Arrival, mono.UnixNano()); isHigh { + r.params.Logger.Debugw( + "high forwarding latency", + "latency", time.Duration(latency), + "queuingLatency", time.Duration(dequeuedAt-pkt.Arrival), + "writeCount", writeCount.Load(), + "isOutOfOrder", pkt.IsOutOfOrder, + "layer", layer, + ) + } + } + + // track video layers + if r.Kind() == webrtc.RTPCodecTypeVideo { + if spatialTrackers[spatialLayer] == nil { + spatialTrackers[spatialLayer] = r.streamTrackerManager.GetTracker(spatialLayer) + if spatialTrackers[spatialLayer] == nil { + if r.videoLayerMode == livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM && pkt.DependencyDescriptor != nil { + r.streamTrackerManager.AddDependencyDescriptorTrackers() + } + spatialTrackers[spatialLayer] = r.streamTrackerManager.AddTracker(spatialLayer) + } + } + if spatialTrackers[spatialLayer] != nil { + spatialTrackers[spatialLayer].Observe( + pkt.Temporal, + len(pkt.RawPacket), + len(pkt.Packet.Payload), + pkt.Packet.Marker, + pkt.Packet.Timestamp, + pkt.DependencyDescriptor, + ) + } + } + + numPacketsForwarded++ + + buffer.ReleaseExtPacket(pkt) + } +} + +func (r *ReceiverBase) DebugInfo() map[string]any { + videoLayerMode := buffer.GetVideoLayerModeForMimeType(r.Mime(), r.TrackInfo()) + info := map[string]any{ + "Mime": r.Mime().String(), + "VideoLayerMode": videoLayerMode.String(), + } + + return info +} + +func (r *ReceiverBase) GetPrimaryReceiverForRed() TrackReceiver { + r.bufferMu.Lock() + defer r.bufferMu.Unlock() + + if !r.isRED || r.IsClosed() { + return r + } + + rt := r.redTransformer.Load() + if rt == nil { + pr := NewRedPrimaryReceiver(r, sfuutils.DownTrackSpreaderParams{ + Threshold: r.lbThreshold, + Logger: r.params.Logger, + }) + r.redTransformer.Store(pr) + return pr + } else { + if pr, ok := rt.(*RedPrimaryReceiver); ok { + return pr + } + } + return nil +} + +func (r *ReceiverBase) GetRedReceiver() TrackReceiver { + r.bufferMu.Lock() + defer r.bufferMu.Unlock() + + if r.isRED || r.IsClosed() { + return r + } + + rt := r.redTransformer.Load() + if rt == nil { + pr := NewRedReceiver(r, sfuutils.DownTrackSpreaderParams{ + Threshold: r.lbThreshold, + Logger: r.params.Logger, + }) + r.redTransformer.Store(pr) + return pr + } else { + if pr, ok := rt.(*RedReceiver); ok { + return pr + } + } + return nil +} + +func (r *ReceiverBase) GetTemporalLayerFpsForSpatial(layer int32) []float32 { + b := r.getBuffer(layer) + if b == nil { + return nil + } + + if r.videoLayerMode != livekit.VideoLayer_MULTIPLE_SPATIAL_LAYERS_PER_STREAM { + return b.GetTemporalLayerFpsForSpatial(0) + } + + return b.GetTemporalLayerFpsForSpatial(layer) +} + +func (r *ReceiverBase) AddOnReady(fn func()) { + // receiver is always ready after created + fn() +} + +func (w *ReceiverBase) handleCodecChange(newCodec webrtc.RTPCodecParameters) { + // codec fallback is not supported mid-session, i.e. change of codec via payload type change, + // set the codec state to invalid once it happens + w.SetCodecState(ReceiverCodecStateInvalid) +} + +func (r *ReceiverBase) AddOnCodecStateChange(f func(webrtc.RTPCodecParameters, ReceiverCodecState)) { + r.codecStateLock.Lock() + r.onCodecStateChange = append(r.onCodecStateChange, f) + r.codecStateLock.Unlock() +} + +func (r *ReceiverBase) CodecState() ReceiverCodecState { + r.codecStateLock.Lock() + defer r.codecStateLock.Unlock() + + return r.codecState +} + +func (r *ReceiverBase) SetCodecState(state ReceiverCodecState) { + r.codecStateLock.Lock() + if r.codecState == state || r.codecState == ReceiverCodecStateInvalid { + r.codecStateLock.Unlock() + return + } + + r.codecState = state + fns := r.onCodecStateChange + r.codecStateLock.Unlock() + + for _, f := range fns { + f(r.params.Codec, state) + } +} + +func (r *ReceiverBase) SetCodecWithState(codec webrtc.RTPCodecParameters, headerExtensions []webrtc.RTPHeaderExtensionParameter, codecState ReceiverCodecState) { + r.checkCodecChanged(codec, headerExtensions) + + r.codecStateLock.Unlock() + if codecState == r.codecState { + r.codecStateLock.Unlock() + return + } + + var fireChange bool + var reason string + onCodecStateChange := r.onCodecStateChange + r.params.Logger.Infow("codec state changed", "from", r.codecState, "to", codecState) + switch codecState { + case ReceiverCodecStateNormal: + // TODO: support codec recovery + r.codecStateLock.Unlock() + return + + case ReceiverCodecStateSuspended: + reason = "codec suspended" + fallthrough + + case ReceiverCodecStateInvalid: + r.codecState = codecState + fireChange = true + reason = "codec invalid" + } + r.codecStateLock.Unlock() + + if fireChange { + r.ClearAllBuffers(reason) + + for _, fn := range onCodecStateChange { + fn(r.params.Codec, codecState) + } + } +} + +func (r *ReceiverBase) checkCodecChanged(codec webrtc.RTPCodecParameters, headerExtensions []webrtc.RTPHeaderExtensionParameter) { + existingFmtp := strings.Split(r.params.Codec.SDPFmtpLine, ";") + slices.Sort(existingFmtp) + checkFmtp := strings.Split(codec.SDPFmtpLine, ";") + slices.Sort(checkFmtp) + if !mime.IsMimeTypeStringEqual(r.params.Codec.MimeType, codec.MimeType) || !slices.Equal(existingFmtp, checkFmtp) || + r.params.Codec.ClockRate != codec.ClockRate { + err := fmt.Errorf("mime: %s -> %s, fmtp: %s -> %s, clockRate: %d -> %d", + r.params.Codec.MimeType, codec.MimeType, + r.params.Codec.SDPFmtpLine, codec.SDPFmtpLine, + r.params.Codec.ClockRate, codec.ClockRate, + ) + r.params.Logger.Errorw("unexpected change in codec", err) + } + + if len(r.params.HeaderExtensions) != len(headerExtensions) { + err := fmt.Errorf("extensions: %d -> %d", len(r.params.HeaderExtensions), len(headerExtensions)) + r.params.Logger.Errorw("unexpected change in extensions length", err) + } +} + +func (r *ReceiverBase) VideoSizes() []buffer.VideoSize { + var sizes []buffer.VideoSize + r.videoSizeMu.RLock() + defer r.videoSizeMu.RUnlock() + for _, v := range r.videoSizes { + if v.Width == 0 || v.Height == 0 { + break + } + sizes = append(sizes, v) + } + + return sizes +} + +func (r *ReceiverBase) OnVideoSizeChanged(f func()) { + r.videoSizeMu.Lock() + r.onVideoSizeChanged = f + r.videoSizeMu.Unlock() +} + +// ----------------------------------------------------------- + +// closes all track senders in parallel, returns when all are closed +func closeTrackSenders(senders []TrackSender) { + wg := sync.WaitGroup{} + for _, dt := range senders { + dt := dt + wg.Add(1) + go func() { + defer wg.Done() + dt.Close() + }() + } + wg.Wait() +} diff --git a/pkg/sfu/redreceiver_test.go b/pkg/sfu/redreceiver_test.go index a8ad32a71..1e3519788 100644 --- a/pkg/sfu/redreceiver_test.go +++ b/pkg/sfu/redreceiver_test.go @@ -42,18 +42,20 @@ func (dt *dummyDowntrack) WriteRTP(p *buffer.ExtPacket, _ int32) int32 { return 1 } -func (dt *dummyDowntrack) TrackInfoAvailable() {} - func TestRedReceiver(t *testing.T) { dt := &dummyDowntrack{TrackSender: &DownTrack{}} t.Run("normal", func(t *testing.T) { w := &WebRTCReceiver{ - isRED: true, - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + isRED: true, + }, } - require.Equal(t, w.GetRedReceiver(), w) + require.Equal(t, w.GetRedReceiver(), w.ReceiverBase) w.isRED = false red := w.GetRedReceiver().(*RedReceiver) require.NotNil(t, red) @@ -75,8 +77,12 @@ func TestRedReceiver(t *testing.T) { t.Run("packet lost and jump", func(t *testing.T) { w := &WebRTCReceiver{ - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + }, } red := w.GetRedReceiver().(*RedReceiver) require.NoError(t, red.AddDownTrack(dt)) @@ -126,8 +132,12 @@ func TestRedReceiver(t *testing.T) { t.Run("unorder and repeat", func(t *testing.T) { w := &WebRTCReceiver{ - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + }, } red := w.GetRedReceiver().(*RedReceiver) require.NoError(t, red.AddDownTrack(dt)) @@ -158,11 +168,15 @@ func TestRedReceiver(t *testing.T) { t.Run("encoding exceed space", func(t *testing.T) { w := &WebRTCReceiver{ - isRED: true, - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + isRED: true, + }, } - require.Equal(t, w.GetRedReceiver(), w) + require.Equal(t, w.GetRedReceiver(), w.ReceiverBase) w.isRED = false red := w.GetRedReceiver().(*RedReceiver) require.NotNil(t, red) @@ -183,11 +197,15 @@ func TestRedReceiver(t *testing.T) { t.Run("large timestamp gap", func(t *testing.T) { w := &WebRTCReceiver{ - isRED: true, - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + isRED: true, + }, } - require.Equal(t, w.GetRedReceiver(), w) + require.Equal(t, w.GetRedReceiver(), w.ReceiverBase) w.isRED = false red := w.GetRedReceiver().(*RedReceiver) require.NotNil(t, red) @@ -283,11 +301,15 @@ func generateRedPkts(t *testing.T, pkts []*rtp.Packet, redCount int) []*rtp.Pack func testRedRedPrimaryReceiver(t *testing.T, maxPktCount, redCount int, sendPktIdx, expectPktIdx []int) { dt := &dummyDowntrack{TrackSender: &DownTrack{}} w := &WebRTCReceiver{ - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), - codec: webrtc.RTPCodecParameters{PayloadType: opusREDPT, RTPCodecCapability: webrtc.RTPCodecCapability{MimeType: "audio/red"}}, + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + Codec: webrtc.RTPCodecParameters{PayloadType: opusREDPT, RTPCodecCapability: webrtc.RTPCodecCapability{MimeType: "audio/red"}}, + }, + }, } - require.Equal(t, w.GetPrimaryReceiverForRed(), w) + require.Equal(t, w.GetPrimaryReceiverForRed(), w.ReceiverBase) w.isRED = true red := w.GetPrimaryReceiverForRed().(*RedPrimaryReceiver) require.NotNil(t, red) @@ -313,10 +335,14 @@ func testRedRedPrimaryReceiver(t *testing.T, maxPktCount, redCount int, sendPktI func TestRedPrimaryReceiver(t *testing.T) { w := &WebRTCReceiver{ - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + }, + }, } - require.Equal(t, w.GetPrimaryReceiverForRed(), w) + require.Equal(t, w.GetPrimaryReceiverForRed(), w.ReceiverBase) w.isRED = true red := w.GetPrimaryReceiverForRed().(*RedPrimaryReceiver) require.NotNil(t, red) @@ -383,11 +409,15 @@ func TestRedPrimaryReceiver(t *testing.T) { t.Run("mixed primary codec", func(t *testing.T) { dt := &dummyDowntrack{TrackSender: &DownTrack{}} w := &WebRTCReceiver{ - kind: webrtc.RTPCodecTypeAudio, - logger: logger.GetLogger(), - codec: webrtc.RTPCodecParameters{PayloadType: opusREDPT, RTPCodecCapability: webrtc.RTPCodecCapability{MimeType: "audio/red"}}, + ReceiverBase: &ReceiverBase{ + params: ReceiverBaseParams{ + Kind: webrtc.RTPCodecTypeAudio, + Logger: logger.GetLogger(), + Codec: webrtc.RTPCodecParameters{PayloadType: opusREDPT, RTPCodecCapability: webrtc.RTPCodecCapability{MimeType: "audio/red"}}, + }, + }, } - require.Equal(t, w.GetPrimaryReceiverForRed(), w) + require.Equal(t, w.GetPrimaryReceiverForRed(), w.ReceiverBase) w.isRED = true red := w.GetPrimaryReceiverForRed().(*RedPrimaryReceiver) require.NotNil(t, red) diff --git a/pkg/sfu/rtpstats/rtpstats_base.go b/pkg/sfu/rtpstats/rtpstats_base.go index 6ddc09f1a..5a21a88ff 100644 --- a/pkg/sfu/rtpstats/rtpstats_base.go +++ b/pkg/sfu/rtpstats/rtpstats_base.go @@ -680,7 +680,7 @@ func (r *rtpStatsBase) updateJitter(ets uint64, packetTime int64) float64 { } func (r *rtpStatsBase) getAndResetSnapshot(snapshotID uint32, extStartSN uint64, extHighestSN uint64) (*snapshot, *snapshot) { - if !r.initialized { + if !r.initialized || snapshotID < cFirstSnapshotID { return nil, nil } diff --git a/pkg/sfu/rtpstats/rtpstats_receiver.go b/pkg/sfu/rtpstats/rtpstats_receiver.go index 5b8cffbbf..7829262bd 100644 --- a/pkg/sfu/rtpstats/rtpstats_receiver.go +++ b/pkg/sfu/rtpstats/rtpstats_receiver.go @@ -751,6 +751,10 @@ func (r *RTPStatsReceiver) MarshalLogObject(e zapcore.ObjectEncoder) error { } func (r *RTPStatsReceiver) ToProto() *livekit.RTPStats { + if r == nil { + return nil + } + r.lock.RLock() defer r.lock.RUnlock() diff --git a/pkg/sfu/rtpstats/rtpstats_sender.go b/pkg/sfu/rtpstats/rtpstats_sender.go index a6194f4cd..186dfba59 100644 --- a/pkg/sfu/rtpstats/rtpstats_sender.go +++ b/pkg/sfu/rtpstats/rtpstats_sender.go @@ -1125,7 +1125,7 @@ func (r *RTPStatsSender) ToProto() *livekit.RTPStats { } func (r *RTPStatsSender) getAndResetSenderSnapshotWindow(senderSnapshotID uint32) (*senderSnapshotWindow, *senderSnapshotWindow) { - if !r.initialized { + if !r.initialized || senderSnapshotID < cFirstSnapshotID { return nil, nil } @@ -1166,7 +1166,7 @@ func (r *RTPStatsSender) getSenderSnapshotWindow(startTime int64) senderSnapshot } func (r *RTPStatsSender) getAndResetSenderSnapshotReceiverView(senderSnapshotID uint32) (*senderSnapshotReceiverView, *senderSnapshotReceiverView) { - if !r.initialized || r.lastRRTime == 0 { + if !r.initialized || r.lastRRTime == 0 || senderSnapshotID < cFirstSnapshotID { return nil, nil } diff --git a/pkg/sfu/track_remote.go b/pkg/sfu/track_remote.go index d868d0f84..b3ba2ae68 100644 --- a/pkg/sfu/track_remote.go +++ b/pkg/sfu/track_remote.go @@ -1,3 +1,17 @@ +// 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 sfu import "github.com/pion/webrtc/v4"