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) +} + +// ------------------------------------------------