From eaf70d5549ac02bc53b725f85ceb5e0de6797d49 Mon Sep 17 00:00:00 2001 From: Raja Subramanian Date: Wed, 28 Jun 2023 13:22:44 +0530 Subject: [PATCH] Pacer in down stream path. (#1835) * Pacer interface to send packets * notify outside lock * use select * use pass through pacer * add error to OnSent * Remove log which could get noisy * Starting TWCC work (#1727) * add packet time * WIP commit * WIP commit * WIP commit * minor comments * Some measurements (#1736) * WIP commit * some notes * WIP commit * variable name change and do not post to closed channel * unlock * clean up * comment * Hooking up some more bits for TWCC (#1752) * wake under lock * Pacer in down stream path. Splitting out only the pacer from a feature branch to introduce the concept of pacer. Currently, there should be no difference in functionality as a pass through pacer is used. Another implementation exists which is just put it in a queue and send it from one goroutine. A potential implementation to try would be data paced by bandwidth estimate. That could include priority queues and such. But, the main goal here is to introduce notion of pacer in the down stream path and prepare for more congestion control possibilities down the line. * Don't need peak detector * remove throttling of write IO errors --- pkg/rtc/mediatracksubscriptions.go | 1 + pkg/rtc/participant.go | 5 + pkg/rtc/transport.go | 32 +- pkg/rtc/transportmanager.go | 5 + pkg/rtc/types/interfaces.go | 3 + .../typesfakes/fake_local_participant.go | 66 +++ pkg/sfu/downtrack.go | 410 +++++++++--------- pkg/sfu/pacer/base.go | 83 ++++ pkg/sfu/pacer/no_queue.go | 80 ++++ pkg/sfu/pacer/pacer.go | 31 ++ pkg/sfu/pacer/packet_time.go | 22 + pkg/sfu/pacer/pass_through.go | 24 + 12 files changed, 537 insertions(+), 225 deletions(-) create mode 100644 pkg/sfu/pacer/base.go create mode 100644 pkg/sfu/pacer/no_queue.go create mode 100644 pkg/sfu/pacer/pacer.go create mode 100644 pkg/sfu/pacer/packet_time.go create mode 100644 pkg/sfu/pacer/pass_through.go diff --git a/pkg/rtc/mediatracksubscriptions.go b/pkg/rtc/mediatracksubscriptions.go index 007428f50..8887176cd 100644 --- a/pkg/rtc/mediatracksubscriptions.go +++ b/pkg/rtc/mediatracksubscriptions.go @@ -104,6 +104,7 @@ func (t *MediaTrackSubscriptions) AddSubscriber(sub types.LocalParticipant, wr * sub.GetBufferFactory(), subscriberID, t.params.ReceiverConfig.PacketBufferSize, + sub.GetPacer(), LoggerWithTrack(sub.GetLogger(), trackID, t.params.IsRelayed), ) if err != nil { diff --git a/pkg/rtc/participant.go b/pkg/rtc/participant.go index 2832fed24..a89315c10 100644 --- a/pkg/rtc/participant.go +++ b/pkg/rtc/participant.go @@ -23,6 +23,7 @@ import ( "github.com/livekit/livekit-server/pkg/sfu" "github.com/livekit/livekit-server/pkg/sfu/buffer" "github.com/livekit/livekit-server/pkg/sfu/connectionquality" + "github.com/livekit/livekit-server/pkg/sfu/pacer" "github.com/livekit/livekit-server/pkg/sfu/streamallocator" "github.com/livekit/livekit-server/pkg/telemetry" "github.com/livekit/livekit-server/pkg/telemetry/prometheus" @@ -231,6 +232,10 @@ func (p *ParticipantImpl) GetAdaptiveStream() bool { return p.params.AdaptiveStream } +func (p *ParticipantImpl) GetPacer() pacer.Pacer { + return p.TransportManager.GetSubscriberPacer() +} + func (p *ParticipantImpl) ID() livekit.ParticipantID { return p.params.SID } diff --git a/pkg/rtc/transport.go b/pkg/rtc/transport.go index 63145bf2c..a164b7eac 100644 --- a/pkg/rtc/transport.go +++ b/pkg/rtc/transport.go @@ -27,6 +27,7 @@ import ( "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/rtc/types" + "github.com/livekit/livekit-server/pkg/sfu/pacer" "github.com/livekit/livekit-server/pkg/sfu/streamallocator" "github.com/livekit/livekit-server/pkg/telemetry" "github.com/livekit/livekit-server/pkg/telemetry/prometheus" @@ -185,6 +186,9 @@ type PCTransport struct { // stream allocator for subscriber PC streamAllocator *streamallocator.StreamAllocator + // only for subscriber PC + pacer pacer.Pacer + previousAnswer *webrtc.SessionDescription // track id -> description map in previous offer sdp previousTrackDescription map[string]*trackDescription @@ -232,9 +236,7 @@ type TransportParams struct { } func newPeerConnection(params TransportParams, onBandwidthEstimator func(estimator cc.BandwidthEstimator)) (*webrtc.PeerConnection, *webrtc.MediaEngine, error) { - directionConfig := params.DirectionConfig - - me, err := createMediaEngine(params.EnabledCodecs, directionConfig) + me, err := createMediaEngine(params.EnabledCodecs, params.DirectionConfig) if err != nil { return nil, nil, err } @@ -305,21 +307,7 @@ func newPeerConnection(params TransportParams, onBandwidthEstimator func(estimat ir := &interceptor.Registry{} if params.IsSendSide { - isSendSideBWE := false - for _, ext := range directionConfig.RTPHeaderExtension.Video { - if ext == sdp.TransportCCURI { - isSendSideBWE = true - break - } - } - for _, ext := range directionConfig.RTPHeaderExtension.Audio { - if ext == sdp.TransportCCURI { - isSendSideBWE = true - break - } - } - - if isSendSideBWE { + if params.CongestionControlConfig.UseSendSideBWE { gf, err := cc.NewInterceptor(func() (cc.BandwidthEstimator, error) { return gcc.NewSendSideBWE( gcc.SendSideBWEInitialBitrate(1*1000*1000), @@ -376,6 +364,7 @@ func NewPCTransport(params TransportParams) (*PCTransport, error) { Logger: params.Logger, }) t.streamAllocator.Start() + t.pacer = pacer.NewPassThrough(params.Logger) } if err := t.createPeerConnection(); err != nil { @@ -414,6 +403,10 @@ func (t *PCTransport) createPeerConnection() error { return nil } +func (t *PCTransport) GetPacer() pacer.Pacer { + return t.pacer +} + func (t *PCTransport) SetSignalingRTT(rtt uint32) { t.signalingRTT.Store(rtt) } @@ -898,6 +891,9 @@ func (t *PCTransport) Close() { if t.streamAllocator != nil { t.streamAllocator.Stop() } + if t.pacer != nil { + t.pacer.Stop() + } _ = t.pc.Close() diff --git a/pkg/rtc/transportmanager.go b/pkg/rtc/transportmanager.go index e34cb3e71..78dd8dee2 100644 --- a/pkg/rtc/transportmanager.go +++ b/pkg/rtc/transportmanager.go @@ -17,6 +17,7 @@ import ( "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/rtc/types" "github.com/livekit/livekit-server/pkg/sfu" + "github.com/livekit/livekit-server/pkg/sfu/pacer" "github.com/livekit/livekit-server/pkg/sfu/streamallocator" "github.com/livekit/livekit-server/pkg/telemetry" "github.com/livekit/protocol/livekit" @@ -283,6 +284,10 @@ func (t *TransportManager) WriteSubscriberRTCP(pkts []rtcp.Packet) error { return t.subscriber.WriteRTCP(pkts) } +func (t *TransportManager) GetSubscriberPacer() pacer.Pacer { + return t.subscriber.GetPacer() +} + func (t *TransportManager) OnPrimaryTransportInitialConnected(f func()) { t.onPrimaryTransportInitialConnected = f } diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index c903d0034..443a0a3d5 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -15,6 +15,7 @@ import ( "github.com/livekit/livekit-server/pkg/routing" "github.com/livekit/livekit-server/pkg/sfu" "github.com/livekit/livekit-server/pkg/sfu/buffer" + "github.com/livekit/livekit-server/pkg/sfu/pacer" ) //go:generate go run github.com/maxbrunsfeld/counterfeiter/v6 -generate @@ -383,6 +384,8 @@ type LocalParticipant interface { // down stream bandwidth management SetSubscriberAllowPause(allowPause bool) SetSubscriberChannelCapacity(channelCapacity int64) + + GetPacer() pacer.Pacer } // Room is a container of participants, and can provide room-level actions diff --git a/pkg/rtc/types/typesfakes/fake_local_participant.go b/pkg/rtc/types/typesfakes/fake_local_participant.go index b4c11ad79..deb07909f 100644 --- a/pkg/rtc/types/typesfakes/fake_local_participant.go +++ b/pkg/rtc/types/typesfakes/fake_local_participant.go @@ -9,6 +9,7 @@ import ( "github.com/livekit/livekit-server/pkg/rtc/types" "github.com/livekit/livekit-server/pkg/sfu" "github.com/livekit/livekit-server/pkg/sfu/buffer" + "github.com/livekit/livekit-server/pkg/sfu/pacer" "github.com/livekit/protocol/auth" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" @@ -252,6 +253,16 @@ type FakeLocalParticipant struct { getLoggerReturnsOnCall map[int]struct { result1 logger.Logger } + GetPacerStub func() pacer.Pacer + getPacerMutex sync.RWMutex + getPacerArgsForCall []struct { + } + getPacerReturns struct { + result1 pacer.Pacer + } + getPacerReturnsOnCall map[int]struct { + result1 pacer.Pacer + } GetPublishedTrackStub func(livekit.TrackID) types.MediaTrack getPublishedTrackMutex sync.RWMutex getPublishedTrackArgsForCall []struct { @@ -2071,6 +2082,59 @@ func (fake *FakeLocalParticipant) GetLoggerReturnsOnCall(i int, result1 logger.L }{result1} } +func (fake *FakeLocalParticipant) GetPacer() pacer.Pacer { + fake.getPacerMutex.Lock() + ret, specificReturn := fake.getPacerReturnsOnCall[len(fake.getPacerArgsForCall)] + fake.getPacerArgsForCall = append(fake.getPacerArgsForCall, struct { + }{}) + stub := fake.GetPacerStub + fakeReturns := fake.getPacerReturns + fake.recordInvocation("GetPacer", []interface{}{}) + fake.getPacerMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeLocalParticipant) GetPacerCallCount() int { + fake.getPacerMutex.RLock() + defer fake.getPacerMutex.RUnlock() + return len(fake.getPacerArgsForCall) +} + +func (fake *FakeLocalParticipant) GetPacerCalls(stub func() pacer.Pacer) { + fake.getPacerMutex.Lock() + defer fake.getPacerMutex.Unlock() + fake.GetPacerStub = stub +} + +func (fake *FakeLocalParticipant) GetPacerReturns(result1 pacer.Pacer) { + fake.getPacerMutex.Lock() + defer fake.getPacerMutex.Unlock() + fake.GetPacerStub = nil + fake.getPacerReturns = struct { + result1 pacer.Pacer + }{result1} +} + +func (fake *FakeLocalParticipant) GetPacerReturnsOnCall(i int, result1 pacer.Pacer) { + fake.getPacerMutex.Lock() + defer fake.getPacerMutex.Unlock() + fake.GetPacerStub = nil + if fake.getPacerReturnsOnCall == nil { + fake.getPacerReturnsOnCall = make(map[int]struct { + result1 pacer.Pacer + }) + } + fake.getPacerReturnsOnCall[i] = struct { + result1 pacer.Pacer + }{result1} +} + func (fake *FakeLocalParticipant) GetPublishedTrack(arg1 livekit.TrackID) types.MediaTrack { fake.getPublishedTrackMutex.Lock() ret, specificReturn := fake.getPublishedTrackReturnsOnCall[len(fake.getPublishedTrackArgsForCall)] @@ -5546,6 +5610,8 @@ func (fake *FakeLocalParticipant) Invocations() map[string][][]interface{} { defer fake.getICEConnectionTypeMutex.RUnlock() fake.getLoggerMutex.RLock() defer fake.getLoggerMutex.RUnlock() + fake.getPacerMutex.RLock() + defer fake.getPacerMutex.RUnlock() fake.getPublishedTrackMutex.RLock() defer fake.getPublishedTrackMutex.RUnlock() fake.getPublishedTracksMutex.RLock() diff --git a/pkg/sfu/downtrack.go b/pkg/sfu/downtrack.go index 848595ae4..b47d9dd33 100644 --- a/pkg/sfu/downtrack.go +++ b/pkg/sfu/downtrack.go @@ -22,6 +22,7 @@ import ( "github.com/livekit/livekit-server/pkg/sfu/buffer" "github.com/livekit/livekit-server/pkg/sfu/connectionquality" dd "github.com/livekit/livekit-server/pkg/sfu/dependencydescriptor" + "github.com/livekit/livekit-server/pkg/sfu/pacer" ) // TrackSender defines an interface send media to remote peer @@ -187,17 +188,17 @@ type DownTrack struct { forwarder *Forwarder - upstreamCodecs []webrtc.RTPCodecParameters - codec webrtc.RTPCodecCapability - rtpHeaderExtensions []webrtc.RTPHeaderExtensionParameter - absSendTimeID int - dependencyDescriptorID int - receiver TrackReceiver - transceiver *webrtc.RTPTransceiver - writeStream webrtc.TrackLocalWriter - rtcpReader *buffer.RTCPReader - onCloseHandler func(willBeResumed bool) - onBinding func(error) + upstreamCodecs []webrtc.RTPCodecParameters + codec webrtc.RTPCodecCapability + absSendTimeExtID int + transportWideExtID int + dependencyDescriptorExtID int + receiver TrackReceiver + transceiver *webrtc.RTPTransceiver + writeStream webrtc.TrackLocalWriter + rtcpReader *buffer.RTCPReader + onCloseHandler func(willBeResumed bool) + onBinding func(error) listenerLock sync.RWMutex receiverReportListeners []ReceiverReportListener @@ -232,6 +233,8 @@ type DownTrack struct { bytesSent atomic.Uint32 bytesRetransmitted atomic.Uint32 + pacer pacer.Pacer + // update stats onStatsUpdate func(dt *DownTrack, stat *livekit.AnalyticsStat) @@ -249,6 +252,7 @@ func NewDownTrack( bf *buffer.Factory, subID livekit.ParticipantID, mt int, + pacer pacer.Pacer, logger logger.Logger, ) (*DownTrack, error) { var kind webrtc.RTPCodecType @@ -272,6 +276,7 @@ func NewDownTrack( upstreamCodecs: codecs, kind: kind, codec: codecs[0].RTPCodecCapability, + pacer: pacer, } d.forwarder = NewForwarder( d.kind, @@ -471,13 +476,14 @@ func (d *DownTrack) SubscriberID() livekit.ParticipantID { return d.subscriberID // Sets RTP header extensions for this track func (d *DownTrack) SetRTPHeaderExtensions(rtpHeaderExtensions []webrtc.RTPHeaderExtensionParameter) { - d.rtpHeaderExtensions = rtpHeaderExtensions for _, ext := range rtpHeaderExtensions { switch ext.URI { case sdp.ABSSendTimeURI: - d.absSendTimeID = ext.ID + d.absSendTimeExtID = ext.ID + case sdp.TransportCCURI: + d.transportWideExtID = ext.ID case dd.ExtensionUrl: - d.dependencyDescriptorID = ext.ID + d.dependencyDescriptorExtID = ext.ID } } } @@ -561,14 +567,6 @@ func (d *DownTrack) keyFrameRequester(generation uint32, layer int32) { // WriteRTP writes an RTP Packet to the DownTrack func (d *DownTrack) WriteRTP(extPkt *buffer.ExtPacket, layer int32) error { - var pool *[]byte - defer func() { - if pool != nil { - PacketFactory.Put(pool) - pool = nil - } - }() - if !d.bound.Load() || !d.connected.Load() { return nil } @@ -581,12 +579,16 @@ func (d *DownTrack) WriteRTP(extPkt *buffer.ExtPacket, layer int32) error { return err } - payload := extPkt.Packet.Payload + var payload []byte + pool := PacketFactory.Get().(*[]byte) if len(tp.codecBytes) != 0 { incomingVP8, _ := extPkt.Payload.(buffer.VP8) - pool = PacketFactory.Get().(*[]byte) payload = d.translateVP8PacketTo(extPkt.Packet, &incomingVP8, tp.codecBytes, pool) } + if payload == nil { + payload = (*pool)[:len(extPkt.Packet.Payload)] + copy(payload, extPkt.Packet.Payload) + } if d.sequencer != nil { d.sequencer.push( @@ -602,45 +604,28 @@ func (d *DownTrack) WriteRTP(extPkt *buffer.ExtPacket, layer int32) error { hdr, err := d.getTranslatedRTPHeader(extPkt, tp) if err != nil { d.logger.Errorw("write rtp packet failed", err) - return err - } - - _, err = d.writeStream.WriteRTP(hdr, payload) - if err != nil { - if !errors.Is(err, io.ErrClosedPipe) { - d.logger.Errorw("write rtp packet failed", err) + if pool != nil { + PacketFactory.Put(pool) } return err } - // STREAM-ALLOCATOR-TODO: remove this stream allocator bytes counter once stream allocator changes fully to pull bytes counter - d.streamAllocatorBytesCounter.Add(uint32(hdr.MarshalSize() + len(payload))) - d.bytesSent.Add(uint32(hdr.MarshalSize() + len(payload))) - - if tp.isSwitchingToMaxSpatial && d.onMaxSubscribedLayerChanged != nil && d.kind == webrtc.RTPCodecTypeVideo { - d.onMaxSubscribedLayerChanged(d, tp.maxSpatialLayer) - } - - if extPkt.KeyFrame { - d.isNACKThrottled.Store(false) - d.rtpStats.UpdateKeyFrame(1) - d.logger.Debugw("forwarding key frame", "layer", layer, "rtpsn", hdr.SequenceNumber, "rtpts", hdr.Timestamp) - } - - if tp.isSwitchingToRequestSpatial { - locked, _ := d.forwarder.CheckSync() - if locked { - d.stopKeyFrameRequester() - } - } - - if tp.isResuming { - if sal := d.getStreamAllocatorListener(); sal != nil { - sal.OnResume(d) - } - } - - d.rtpStats.Update(hdr, len(payload), 0, extPkt.Arrival) + d.pacer.Enqueue(pacer.Packet{ + Header: hdr, + Extensions: []pacer.ExtensionData{{ID: uint8(d.dependencyDescriptorExtID), Payload: tp.ddBytes}}, + Payload: payload, + AbsSendTimeExtID: uint8(d.absSendTimeExtID), + TransportWideExtID: uint8(d.transportWideExtID), + WriteStream: d.writeStream, + Metadata: sendPacketMetadata{ + layer: layer, + arrival: extPkt.Arrival, + isKeyFrame: extPkt.KeyFrame, + tp: tp, + pool: pool, + }, + OnSent: d.packetSent, + }) return nil } @@ -704,23 +689,23 @@ func (d *DownTrack) WritePaddingRTP(bytesToSend int, paddingOnMute bool, forceMa CSRC: []uint32{}, } - err = d.writeRTPHeaderExtensions(&hdr) - if err != nil { - return bytesSent - } - payload := make([]byte, RTPPaddingMaxPayloadSize) // last byte of padding has padding size including that byte payload[RTPPaddingMaxPayloadSize-1] = byte(RTPPaddingMaxPayloadSize) - _, err = d.writeStream.WriteRTP(&hdr, payload) - if err != nil { - return bytesSent - } - - if !paddingOnMute { - d.rtpStats.Update(&hdr, 0, len(payload), time.Now()) - } + d.pacer.Enqueue(pacer.Packet{ + Header: &hdr, + Payload: payload, + AbsSendTimeExtID: uint8(d.absSendTimeExtID), + TransportWideExtID: uint8(d.transportWideExtID), + WriteStream: d.writeStream, + Metadata: sendPacketMetadata{ + isPadding: true, + disableCounter: true, + disableRTPStats: paddingOnMute, + }, + OnSent: d.packetSent, + }) // // Register with sequencer with invalid layer so that NACKs for these can be filtered out. @@ -734,6 +719,7 @@ func (d *DownTrack) WritePaddingRTP(bytesToSend int, paddingOnMute bool, forceMa bytesSent += hdr.MarshalSize() + len(payload) } + // STREAM_ALLOCATOR-TODO: change this to pull this counter from stream allocator so that counter can be update in pacer callback return bytesSent } @@ -1123,16 +1109,16 @@ func (d *DownTrack) writeBlankFrameRTP(duration float32, generation uint32) chan return } - var writeBlankFrame func(*rtp.Header, bool) (int, error) + var getBlankFrame func(bool) ([]byte, error) switch d.mime { case "audio/opus": - writeBlankFrame = d.writeOpusBlankFrame + getBlankFrame = d.getOpusBlankFrame case "audio/red": - writeBlankFrame = d.writeOpusRedBlankFrame + getBlankFrame = d.getOpusRedBlankFrame case "video/vp8": - writeBlankFrame = d.writeVP8BlankFrame + getBlankFrame = d.getVP8BlankFrame case "video/h264": - writeBlankFrame = d.writeH264BlankFrame + getBlankFrame = d.getH264BlankFrame default: close(done) return @@ -1177,24 +1163,24 @@ func (d *DownTrack) writeBlankFrameRTP(duration float32, generation uint32) chan CSRC: []uint32{}, } - err = d.writeRTPHeaderExtensions(&hdr) + payload, err := getBlankFrame(frameEndNeeded) if err != nil { - d.logger.Warnw("could not write header extension for blank frame", err) + d.logger.Warnw("could not get blank frame", err) close(done) return } - pktSize, err := writeBlankFrame(&hdr, frameEndNeeded) - if err != nil { - if err != io.ErrClosedPipe { - d.logger.Warnw("could not write blank frame", err) - } - close(done) - return - } - - d.streamAllocatorBytesCounter.Add(uint32(pktSize)) - d.bytesSent.Add(uint32(pktSize)) + d.pacer.Enqueue(pacer.Packet{ + Header: &hdr, + Payload: payload, + AbsSendTimeExtID: uint8(d.absSendTimeExtID), + TransportWideExtID: uint8(d.transportWideExtID), + WriteStream: d.writeStream, + Metadata: sendPacketMetadata{ + isBlankFrame: true, + }, + OnSent: d.packetSent, + }) // only the first frame will need frameEndNeeded to close out the // previous picture, rest are small key frames (for the video case) @@ -1209,22 +1195,17 @@ func (d *DownTrack) writeBlankFrameRTP(duration float32, generation uint32) chan return done } -func (d *DownTrack) writeOpusBlankFrame(hdr *rtp.Header, frameEndNeeded bool) (int, error) { +func (d *DownTrack) getOpusBlankFrame(_frameEndNeeded bool) ([]byte, error) { // silence frame // Used shortly after muting to ensure residual noise does not keep // generating noise at the decoder after the stream is stopped // i. e. comfort noise generation actually not producing something comfortable. payload := make([]byte, len(OpusSilenceFrame)) copy(payload[0:], OpusSilenceFrame) - - _, err := d.writeStream.WriteRTP(hdr, payload) - if err == nil { - d.rtpStats.Update(hdr, len(payload), 0, time.Now()) - } - return hdr.MarshalSize() + len(payload), err + return payload, nil } -func (d *DownTrack) writeOpusRedBlankFrame(hdr *rtp.Header, frameEndNeeded bool) (int, error) { +func (d *DownTrack) getOpusRedBlankFrame(_frameEndNeeded bool) ([]byte, error) { // primary only silence frame for opus/red, there is no need to contain redundant silent frames payload := make([]byte, len(OpusSilenceFrame)+1) @@ -1235,18 +1216,13 @@ func (d *DownTrack) writeOpusRedBlankFrame(hdr *rtp.Header, frameEndNeeded bool) // +-+-+-+-+-+-+-+-+ payload[0] = opusPT copy(payload[1:], OpusSilenceFrame) - - _, err := d.writeStream.WriteRTP(hdr, payload) - if err == nil { - d.rtpStats.Update(hdr, len(payload), 0, time.Now()) - } - return hdr.MarshalSize() + len(payload), err + return payload, nil } -func (d *DownTrack) writeVP8BlankFrame(hdr *rtp.Header, frameEndNeeded bool) (int, error) { +func (d *DownTrack) getVP8BlankFrame(frameEndNeeded bool) ([]byte, error) { blankVP8, err := d.forwarder.GetPadding(frameEndNeeded) if err != nil { - return 0, err + return nil, err } // 8x8 key frame @@ -1256,15 +1232,10 @@ func (d *DownTrack) writeVP8BlankFrame(hdr *rtp.Header, frameEndNeeded bool) (in payload := make([]byte, len(blankVP8)+len(VP8KeyFrame8x8)) copy(payload[:len(blankVP8)], blankVP8) copy(payload[len(blankVP8):], VP8KeyFrame8x8) - - _, err = d.writeStream.WriteRTP(hdr, payload) - if err == nil { - d.rtpStats.Update(hdr, len(payload), 0, time.Now()) - } - return hdr.MarshalSize() + len(payload), err + return payload, nil } -func (d *DownTrack) writeH264BlankFrame(hdr *rtp.Header, frameEndNeeded bool) (int, error) { +func (d *DownTrack) getH264BlankFrame(_frameEndNeeded bool) ([]byte, error) { // TODO - Jie Zeng // now use STAP-A to compose sps, pps, idr together, most decoder support packetization-mode 1. // if client only support packetization-mode 0, use single nalu unit packet @@ -1279,11 +1250,7 @@ func (d *DownTrack) writeH264BlankFrame(hdr *rtp.Header, frameEndNeeded bool) (i offset += len(payload) } payload := buf[:offset] - _, err := d.writeStream.WriteRTP(hdr, payload) - if err == nil { - d.rtpStats.Update(hdr, len(payload), 0, time.Now()) - } - return hdr.MarshalSize() + offset, err + return payload, nil } func (d *DownTrack) handleRTCP(bytes []byte) { @@ -1416,14 +1383,6 @@ func (d *DownTrack) retransmitPackets(nacks []uint16) { return } - var pool *[]byte - defer func() { - if pool != nil { - PacketFactory.Put(pool) - pool = nil - } - }() - src := PacketFactory.Get().(*[]byte) defer PacketFactory.Put(src) @@ -1443,11 +1402,6 @@ func (d *DownTrack) retransmitPackets(nacks []uint16) { Attempts: meta.nacked, }) - if pool != nil { - PacketFactory.Put(pool) - pool = nil - } - pktBuff := *src n, err := d.receiver.ReadRTP(pktBuff, uint8(meta.layer), meta.sourceSeqNo) if err != nil { @@ -1471,41 +1425,38 @@ func (d *DownTrack) retransmitPackets(nacks []uint16) { pkt.Header.SSRC = d.ssrc pkt.Header.PayloadType = d.payloadType - payload := pkt.Payload + var payload []byte + pool := PacketFactory.Get().(*[]byte) if d.mime == "video/vp8" && len(pkt.Payload) > 0 { var incomingVP8 buffer.VP8 if err = incomingVP8.Unmarshal(pkt.Payload); err != nil { d.logger.Errorw("unmarshalling VP8 packet err", err) + PacketFactory.Put(pool) continue } if len(meta.codecBytes) != 0 { - pool = PacketFactory.Get().(*[]byte) payload = d.translateVP8PacketTo(&pkt, &incomingVP8, meta.codecBytes, pool) } } - - var extraExtensions []extensionData - if d.dependencyDescriptorID != 0 && len(meta.ddBytes) != 0 { - extraExtensions = append(extraExtensions, extensionData{ - id: uint8(d.dependencyDescriptorID), - payload: meta.ddBytes, - }) - } - err = d.writeRTPHeaderExtensions(&pkt.Header, extraExtensions...) - if err != nil { - d.logger.Errorw("writing rtp header extensions err", err) - continue + if payload == nil { + payload = (*pool)[:len(pkt.Payload)] + copy(payload, pkt.Payload) } - if _, err = d.writeStream.WriteRTP(&pkt.Header, payload); err != nil { - d.logger.Errorw("writing rtx packet err", err) - } else { - d.streamAllocatorBytesCounter.Add(uint32(pkt.Header.MarshalSize() + len(payload))) - d.bytesRetransmitted.Add(uint32(pkt.Header.MarshalSize() + len(payload))) - - d.rtpStats.Update(&pkt.Header, len(payload), 0, time.Now()) - } + d.pacer.Enqueue(pacer.Packet{ + Header: &pkt.Header, + Extensions: []pacer.ExtensionData{{ID: uint8(d.dependencyDescriptorExtID), Payload: meta.ddBytes}}, + Payload: payload, + AbsSendTimeExtID: uint8(d.absSendTimeExtID), + TransportWideExtID: uint8(d.transportWideExtID), + WriteStream: d.writeStream, + Metadata: sendPacketMetadata{ + isRTX: true, + pool: pool, + }, + OnSent: d.packetSent, + }) } d.totalRepeatedNACKs.Add(numRepeatedNACKs) @@ -1528,38 +1479,6 @@ func (d *DownTrack) retransmitPackets(nacks []uint16) { } } -type extensionData struct { - id uint8 - payload []byte -} - -// writes RTP header extensions of track -func (d *DownTrack) writeRTPHeaderExtensions(hdr *rtp.Header, extraExtensions ...extensionData) error { - // clear out extensions that may have been in the forwarded header - hdr.Extension = false - hdr.ExtensionProfile = 0 - hdr.Extensions = []rtp.Extension{} - - for _, ext := range extraExtensions { - hdr.SetExtension(ext.id, ext.payload) - } - - if d.absSendTimeID != 0 { - sendTime := rtp.NewAbsSendTimeExtension(time.Now()) - b, err := sendTime.Marshal() - if err != nil { - return err - } - - err = hdr.SetExtension(uint8(d.absSendTimeID), b) - if err != nil { - return err - } - } - - return nil -} - func (d *DownTrack) getTranslatedRTPHeader(extPkt *buffer.ExtPacket, tp *TranslationParams) (*rtp.Header, error) { tpRTP := tp.rtp hdr := extPkt.Packet.Header @@ -1571,18 +1490,6 @@ func (d *DownTrack) getTranslatedRTPHeader(extPkt *buffer.ExtPacket, tp *Transla hdr.Marker = tp.marker } - var extension []extensionData - if d.dependencyDescriptorID != 0 && len(tp.ddBytes) != 0 { - extension = append(extension, extensionData{ - id: uint8(d.dependencyDescriptorID), - payload: tp.ddBytes, - }) - } - err := d.writeRTPHeaderExtensions(&hdr, extension...) - if err != nil { - return nil, err - } - return &hdr, nil } @@ -1739,20 +1646,24 @@ func (d *DownTrack) sendSilentFrameOnMuteForOpus() { CSRC: []uint32{}, } - err = d.writeRTPHeaderExtensions(&hdr) + payload, err := d.getOpusBlankFrame(false) if err != nil { - d.logger.Warnw("could not write header extension for blank frame", err) + d.logger.Warnw("could not get blank frame", err) return } - payload := make([]byte, len(OpusSilenceFrame)) - copy(payload[0:], OpusSilenceFrame) - - _, err := d.writeStream.WriteRTP(&hdr, payload) - if err != nil { - d.logger.Warnw("could not write blank frame", err) - return - } + d.pacer.Enqueue(pacer.Packet{ + Header: &hdr, + Payload: payload, + AbsSendTimeExtID: uint8(d.absSendTimeExtID), + TransportWideExtID: uint8(d.transportWideExtID), + WriteStream: d.writeStream, + Metadata: sendPacketMetadata{ + isBlankFrame: true, + disableRTPStats: true, + }, + OnSent: d.packetSent, + }) } numFrames-- @@ -1763,3 +1674,88 @@ func (d *DownTrack) sendSilentFrameOnMuteForOpus() { func (d *DownTrack) HandleRTCPSenderReportData(_payloadType webrtc.PayloadType, _layer int32, _srData *buffer.RTCPSenderReportData) error { return nil } + +type sendPacketMetadata struct { + layer int32 + arrival time.Time + isKeyFrame bool + isRTX bool + isPadding bool + isBlankFrame bool + disableCounter bool + disableRTPStats bool + tp *TranslationParams + pool *[]byte +} + +func (d *DownTrack) packetSent(md interface{}, hdr *rtp.Header, payloadSize int, sendTime time.Time, sendError error) { + spmd, ok := md.(sendPacketMetadata) + if !ok { + d.logger.Errorw("invalid send packet metadata", nil) + return + } + + if spmd.pool != nil { + PacketFactory.Put(spmd.pool) + } + + if sendError != nil { + return + } + + headerSize := hdr.MarshalSize() + if !spmd.disableCounter { + // STREAM-ALLOCATOR-TODO: remove this stream allocator bytes counter once stream allocator changes fully to pull bytes counter + size := uint32(headerSize + payloadSize) + d.streamAllocatorBytesCounter.Add(size) + if spmd.isRTX { + d.bytesRetransmitted.Add(size) + } else { + d.bytesSent.Add(size) + } + } + + if !spmd.disableRTPStats { + packetTime := spmd.arrival + if packetTime.IsZero() { + packetTime = sendTime + } + if spmd.isPadding { + d.rtpStats.Update(hdr, 0, payloadSize, packetTime) + } else { + d.rtpStats.Update(hdr, payloadSize, 0, packetTime) + } + } + + if spmd.isKeyFrame { + d.isNACKThrottled.Store(false) + d.rtpStats.UpdateKeyFrame(1) + d.logger.Debugw( + "forwarding key frame", + "layer", spmd.layer, + "rtpsn", hdr.SequenceNumber, + "rtpts", hdr.Timestamp, + ) + } + + if spmd.tp != nil { + if spmd.tp.isSwitchingToMaxSpatial && d.onMaxSubscribedLayerChanged != nil && d.kind == webrtc.RTPCodecTypeVideo { + d.onMaxSubscribedLayerChanged(d, spmd.tp.maxSpatialLayer) + } + + if spmd.tp.isSwitchingToRequestSpatial { + locked, _ := d.forwarder.CheckSync() + if locked { + d.stopKeyFrameRequester() + } + } + + if spmd.tp.isResuming { + if sal := d.getStreamAllocatorListener(); sal != nil { + sal.OnResume(d) + } + } + } +} + +// ------------------------------------------------------------------------------- diff --git a/pkg/sfu/pacer/base.go b/pkg/sfu/pacer/base.go new file mode 100644 index 000000000..fef7b413e --- /dev/null +++ b/pkg/sfu/pacer/base.go @@ -0,0 +1,83 @@ +package pacer + +import ( + "errors" + "io" + "time" + + "github.com/livekit/protocol/logger" + "github.com/pion/rtp" +) + +type Base struct { + logger logger.Logger + + packetTime *PacketTime +} + +func NewBase(logger logger.Logger) *Base { + return &Base{ + logger: logger, + packetTime: NewPacketTime(), + } +} + +func (b *Base) SendPacket(p *Packet) error { + var sendingAt time.Time + var err error + defer func() { + if p.OnSent != nil { + p.OnSent(p.Metadata, p.Header, len(p.Payload), sendingAt, err) + } + }() + + sendingAt, err = b.writeRTPHeaderExtensions(p) + if err != nil { + b.logger.Errorw("writing rtp header extensions err", err) + return err + } + + _, err = p.WriteStream.WriteRTP(p.Header, p.Payload) + if err != nil { + if !errors.Is(err, io.ErrClosedPipe) { + b.logger.Errorw("write rtp packet failed", err) + } + return err + } + + return nil +} + +// writes RTP header extensions of track +func (b *Base) writeRTPHeaderExtensions(p *Packet) (time.Time, error) { + // clear out extensions that may have been in the forwarded header + p.Header.Extension = false + p.Header.ExtensionProfile = 0 + p.Header.Extensions = []rtp.Extension{} + + for _, ext := range p.Extensions { + if ext.ID == 0 || len(ext.Payload) == 0 { + continue + } + + p.Header.SetExtension(ext.ID, ext.Payload) + } + + sendingAt := b.packetTime.Get() + if p.AbsSendTimeExtID != 0 { + sendTime := rtp.NewAbsSendTimeExtension(sendingAt) + b, err := sendTime.Marshal() + if err != nil { + return time.Time{}, err + } + + err = p.Header.SetExtension(p.AbsSendTimeExtID, b) + if err != nil { + return time.Time{}, err + } + } + + return sendingAt, nil +} + +// ------------------------------------------------ diff --git a/pkg/sfu/pacer/no_queue.go b/pkg/sfu/pacer/no_queue.go new file mode 100644 index 000000000..b34b994ae --- /dev/null +++ b/pkg/sfu/pacer/no_queue.go @@ -0,0 +1,80 @@ +package pacer + +import ( + "sync" + + "github.com/gammazero/deque" + "github.com/livekit/protocol/logger" +) + +type NoQueue struct { + *Base + + logger logger.Logger + + lock sync.RWMutex + packets deque.Deque[Packet] + wake chan struct{} + isStopped bool +} + +func NewNoQueue(logger logger.Logger) *NoQueue { + n := &NoQueue{ + Base: NewBase(logger), + logger: logger, + wake: make(chan struct{}, 1), + } + n.packets.SetMinCapacity(9) + + go n.sendWorker() + return n +} + +func (n *NoQueue) Stop() { + n.lock.Lock() + if n.isStopped { + n.lock.Unlock() + return + } + + close(n.wake) + n.isStopped = true + n.lock.Unlock() +} + +func (n *NoQueue) Enqueue(p Packet) { + n.lock.Lock() + defer n.lock.Unlock() + + n.packets.PushBack(p) + if n.packets.Len() == 1 && !n.isStopped { + select { + case n.wake <- struct{}{}: + default: + } + } +} + +func (n *NoQueue) sendWorker() { + for { + <-n.wake + for { + n.lock.Lock() + if n.isStopped { + n.lock.Unlock() + return + } + + if n.packets.Len() == 0 { + n.lock.Unlock() + break + } + p := n.packets.PopFront() + n.lock.Unlock() + + n.Base.SendPacket(&p) + } + } +} + +// ------------------------------------------------ diff --git a/pkg/sfu/pacer/pacer.go b/pkg/sfu/pacer/pacer.go new file mode 100644 index 000000000..3be8a8a86 --- /dev/null +++ b/pkg/sfu/pacer/pacer.go @@ -0,0 +1,31 @@ +package pacer + +import ( + "time" + + "github.com/pion/rtp" + "github.com/pion/webrtc/v3" +) + +type ExtensionData struct { + ID uint8 + Payload []byte +} + +type Packet struct { + Header *rtp.Header + Extensions []ExtensionData + Payload []byte + AbsSendTimeExtID uint8 + TransportWideExtID uint8 + WriteStream webrtc.TrackLocalWriter + Metadata interface{} + OnSent func(md interface{}, sentHeader *rtp.Header, payloadSize int, sentTime time.Time, sendError error) +} + +type Pacer interface { + Enqueue(p Packet) + Stop() +} + +// ------------------------------------------------ diff --git a/pkg/sfu/pacer/packet_time.go b/pkg/sfu/pacer/packet_time.go new file mode 100644 index 000000000..3dce57e3a --- /dev/null +++ b/pkg/sfu/pacer/packet_time.go @@ -0,0 +1,22 @@ +package pacer + +import ( + "time" +) + +type PacketTime struct { + baseTime time.Time +} + +func NewPacketTime() *PacketTime { + return &PacketTime{ + baseTime: time.Now(), + } +} + +func (p *PacketTime) Get() time.Time { + // construct current time based on monotonic clock + return p.baseTime.Add(time.Since(p.baseTime)) +} + +// ------------------------------------------------ diff --git a/pkg/sfu/pacer/pass_through.go b/pkg/sfu/pacer/pass_through.go new file mode 100644 index 000000000..ccbefbd61 --- /dev/null +++ b/pkg/sfu/pacer/pass_through.go @@ -0,0 +1,24 @@ +package pacer + +import ( + "github.com/livekit/protocol/logger" +) + +type PassThrough struct { + *Base +} + +func NewPassThrough(logger logger.Logger) *PassThrough { + return &PassThrough{ + Base: NewBase(logger), + } +} + +func (p *PassThrough) Stop() { +} + +func (p *PassThrough) Enqueue(pkt Packet) { + p.Base.SendPacket(&pkt) +} + +// ------------------------------------------------