Separate from ion-sfu (#171)

* Separate from ion-sfu

changes:
1. extract pkg/buffer, twcc, sfu, relay, stats, logger

2. to solve cycle import, move ion-sfu/pkg/logger to pkg/sfu/logger

3. replace pion/ion-sfu => ./
reason: will change import pion/ion-sfu/pkg/* to livekit-server/pkg/*
after this pr merged. Just not change any code in this pr, because it
will confused with the separate code from ion-sfu in review.

* Move code from ion-sfu to pkg/sfu

* fix build error for resovle conflict

Co-authored-by: cnderrauber <zengjie9004@gmail.com>
This commit is contained in:
cnderrauber
2021-11-09 12:03:16 +08:00
committed by GitHub
co-authored by cnderrauber
parent 289ebd32ff
commit 1e1aaeb86b
58 changed files with 12210 additions and 384 deletions
+114
View File
@@ -0,0 +1,114 @@
package buffer
import (
"encoding/binary"
"math"
)
const maxPktSize = 1500
type Bucket struct {
buf []byte
src *[]byte
init bool
step int
headSN uint16
maxSteps int
}
func NewBucket(buf *[]byte) *Bucket {
return &Bucket{
src: buf,
buf: *buf,
maxSteps: int(math.Floor(float64(len(*buf))/float64(maxPktSize))) - 1,
}
}
func (b *Bucket) AddPacket(pkt []byte, sn uint16, latest bool) ([]byte, error) {
if !b.init {
b.headSN = sn - 1
b.init = true
}
if !latest {
return b.set(sn, pkt)
}
diff := sn - b.headSN
b.headSN = sn
for i := uint16(1); i < diff; i++ {
b.step++
if b.step >= b.maxSteps {
b.step = 0
}
}
return b.push(pkt), nil
}
func (b *Bucket) GetPacket(buf []byte, sn uint16) (i int, err error) {
p := b.get(sn)
if p == nil {
err = errPacketNotFound
return
}
i = len(p)
if cap(buf) < i {
err = errBufferTooSmall
return
}
if len(buf) < i {
buf = buf[:i]
}
copy(buf, p)
return
}
func (b *Bucket) push(pkt []byte) []byte {
binary.BigEndian.PutUint16(b.buf[b.step*maxPktSize:], uint16(len(pkt)))
off := b.step*maxPktSize + 2
copy(b.buf[off:], pkt)
b.step++
if b.step > b.maxSteps {
b.step = 0
}
return b.buf[off : off+len(pkt)]
}
func (b *Bucket) get(sn uint16) []byte {
pos := b.step - int(b.headSN-sn+1)
if pos < 0 {
if pos*-1 > b.maxSteps+1 {
return nil
}
pos = b.maxSteps + pos + 1
}
off := pos * maxPktSize
if off > len(b.buf) {
return nil
}
if binary.BigEndian.Uint16(b.buf[off+4:off+6]) != sn {
return nil
}
sz := int(binary.BigEndian.Uint16(b.buf[off : off+2]))
return b.buf[off+2 : off+2+sz]
}
func (b *Bucket) set(sn uint16, pkt []byte) ([]byte, error) {
if b.headSN-sn >= uint16(b.maxSteps+1) {
return nil, errPacketTooOld
}
pos := b.step - int(b.headSN-sn+1)
if pos < 0 {
pos = b.maxSteps + pos + 1
}
off := pos * maxPktSize
if off > len(b.buf) || off < 0 {
return nil, errPacketTooOld
}
// Do not overwrite if packet exist
if binary.BigEndian.Uint16(b.buf[off+4:off+6]) == sn {
return nil, errRTXPacket
}
binary.BigEndian.PutUint16(b.buf[off:], uint16(len(pkt)))
copy(b.buf[off+2:], pkt)
return b.buf[off+2 : off+2+len(pkt)], nil
}
+139
View File
@@ -0,0 +1,139 @@
package buffer
import (
"testing"
"github.com/pion/rtp"
"github.com/stretchr/testify/assert"
)
var TestPackets = []*rtp.Packet{
{
Header: rtp.Header{
SequenceNumber: 1,
},
},
{
Header: rtp.Header{
SequenceNumber: 3,
},
},
{
Header: rtp.Header{
SequenceNumber: 4,
},
},
{
Header: rtp.Header{
SequenceNumber: 6,
},
},
{
Header: rtp.Header{
SequenceNumber: 7,
},
},
{
Header: rtp.Header{
SequenceNumber: 10,
},
},
}
func Test_queue(t *testing.T) {
b := make([]byte, 25000)
q := NewBucket(&b)
for _, p := range TestPackets {
p := p
buf, err := p.Marshal()
assert.NoError(t, err)
assert.NotPanics(t, func() {
q.AddPacket(buf, p.SequenceNumber, true)
})
}
var expectedSN uint16
expectedSN = 6
np := rtp.Packet{}
buff := make([]byte, maxPktSize)
i, err := q.GetPacket(buff, 6)
assert.NoError(t, err)
err = np.Unmarshal(buff[:i])
assert.NoError(t, err)
assert.Equal(t, expectedSN, np.SequenceNumber)
np2 := &rtp.Packet{
Header: rtp.Header{
SequenceNumber: 8,
},
}
buf, err := np2.Marshal()
assert.NoError(t, err)
expectedSN = 8
q.AddPacket(buf, 8, false)
i, err = q.GetPacket(buff, expectedSN)
assert.NoError(t, err)
err = np.Unmarshal(buff[:i])
assert.NoError(t, err)
assert.Equal(t, expectedSN, np.SequenceNumber)
_, err = q.AddPacket(buf, 8, false)
assert.ErrorIs(t, err, errRTXPacket)
}
func Test_queue_edges(t *testing.T) {
var TestPackets = []*rtp.Packet{
{
Header: rtp.Header{
SequenceNumber: 65533,
},
},
{
Header: rtp.Header{
SequenceNumber: 65534,
},
},
{
Header: rtp.Header{
SequenceNumber: 2,
},
},
}
b := make([]byte, 25000)
q := NewBucket(&b)
for _, p := range TestPackets {
p := p
assert.NotNil(t, p)
assert.NotPanics(t, func() {
p := p
buf, err := p.Marshal()
assert.NoError(t, err)
assert.NotPanics(t, func() {
q.AddPacket(buf, p.SequenceNumber, true)
})
})
}
var expectedSN uint16
expectedSN = 65534
np := rtp.Packet{}
buff := make([]byte, maxPktSize)
i, err := q.GetPacket(buff, expectedSN)
assert.NoError(t, err)
err = np.Unmarshal(buff[:i])
assert.NoError(t, err)
assert.Equal(t, expectedSN, np.SequenceNumber)
np2 := rtp.Packet{
Header: rtp.Header{
SequenceNumber: 65535,
},
}
buf, err := np2.Marshal()
assert.NoError(t, err)
q.AddPacket(buf, np2.SequenceNumber, false)
i, err = q.GetPacket(buff, expectedSN+1)
assert.NoError(t, err)
err = np.Unmarshal(buff[:i])
assert.NoError(t, err)
assert.Equal(t, expectedSN+1, np.SequenceNumber)
}
+715
View File
@@ -0,0 +1,715 @@
package buffer
import (
"encoding/binary"
"io"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/gammazero/deque"
"github.com/go-logr/logr"
"github.com/pion/rtcp"
"github.com/pion/rtp"
"github.com/pion/sdp/v3"
"github.com/pion/webrtc/v3"
)
const (
MaxSN = 1 << 16
reportDelta = 1e9
)
// Logger is an implementation of logr.Logger. If is not provided - will be turned off.
var Logger logr.Logger = logr.Discard()
type pendingPackets struct {
arrivalTime int64
packet []byte
}
type ExtPacket struct {
Head bool
Cycle uint32
Arrival int64
Packet rtp.Packet
Payload interface{}
KeyFrame bool
}
// Buffer contains all packets
type Buffer struct {
sync.Mutex
bucket *Bucket
nacker *nackQueue
videoPool *sync.Pool
audioPool *sync.Pool
codecType webrtc.RTPCodecType
extPackets deque.Deque
pPackets []pendingPackets
closeOnce sync.Once
mediaSSRC uint32
clockRate uint32
maxBitrate uint64
lastReport int64
twccExt uint8
audioExt uint8
bound bool
closed atomicBool
mime string
// supported feedbacks
remb bool
nack bool
twcc bool
audioLevel bool
minPacketProbe int
lastPacketRead int
maxTemporalLayer int32
bitrate atomic.Value
bitrateHelper [4]uint64
lastSRNTPTime uint64
lastSRRTPTime uint32
lastSRRecv int64 // Represents wall clock of the most recent sender report arrival
baseSN uint16
lastRtcpPacketTime int64 // Time the last RTCP packet was received.
lastRtcpSrTime int64 // Time the last RTCP SR was received. Required for DLSR computation.
lastTransit uint32
seqHdlr SeqWrapHandler
stats Stats
latestTimestamp uint32 // latest received RTP timestamp on packet
latestTimestampTime int64 // Time of the latest timestamp (in nanos since unix epoch)
lastFractionLostToReport uint8 // Last fractionlost from subscribers, should report to publisher; Audio only
// callbacks
onClose func()
onAudioLevel func(level uint8)
feedbackCB func([]rtcp.Packet)
feedbackTWCC func(sn uint16, timeNS int64, marker bool)
// logger
logger logr.Logger
}
type Stats struct {
LastExpected uint32
LastReceived uint32
LostRate float32
PacketCount uint32 // Number of packets received from this source.
Jitter float64 // An estimate of the statistical variance of the RTP data packet inter-arrival time.
TotalByte uint64
}
// BufferOptions provides configuration options for the buffer
type Options struct {
MaxBitRate uint64
}
// NewBuffer constructs a new Buffer
func NewBuffer(ssrc uint32, vp, ap *sync.Pool, logger logr.Logger) *Buffer {
b := &Buffer{
mediaSSRC: ssrc,
videoPool: vp,
audioPool: ap,
logger: logger,
}
b.bitrate.Store(make([]uint64, len(b.bitrateHelper)))
b.extPackets.SetMinCapacity(7)
return b
}
func (b *Buffer) Bind(params webrtc.RTPParameters, o Options) {
b.Lock()
defer b.Unlock()
codec := params.Codecs[0]
b.clockRate = codec.ClockRate
b.maxBitrate = o.MaxBitRate
b.mime = strings.ToLower(codec.MimeType)
switch {
case strings.HasPrefix(b.mime, "audio/"):
b.codecType = webrtc.RTPCodecTypeAudio
b.bucket = NewBucket(b.audioPool.Get().(*[]byte))
case strings.HasPrefix(b.mime, "video/"):
b.codecType = webrtc.RTPCodecTypeVideo
b.bucket = NewBucket(b.videoPool.Get().(*[]byte))
default:
b.codecType = webrtc.RTPCodecType(0)
}
for _, ext := range params.HeaderExtensions {
if ext.URI == sdp.TransportCCURI {
b.twccExt = uint8(ext.ID)
break
}
}
if b.codecType == webrtc.RTPCodecTypeVideo {
for _, fb := range codec.RTCPFeedback {
switch fb.Type {
case webrtc.TypeRTCPFBGoogREMB:
b.logger.V(1).Info("Setting feedback", "type", "webrtc.TypeRTCPFBGoogREMB")
b.remb = true
case webrtc.TypeRTCPFBTransportCC:
b.logger.V(1).Info("Setting feedback", "type", webrtc.TypeRTCPFBTransportCC)
b.twcc = true
case webrtc.TypeRTCPFBNACK:
b.logger.V(1).Info("Setting feedback", "type", webrtc.TypeRTCPFBNACK)
b.nacker = newNACKQueue()
b.nack = true
}
}
} else if b.codecType == webrtc.RTPCodecTypeAudio {
for _, h := range params.HeaderExtensions {
if h.URI == sdp.AudioLevelURI {
b.audioLevel = true
b.audioExt = uint8(h.ID)
}
}
}
for _, pp := range b.pPackets {
b.calc(pp.packet, pp.arrivalTime)
}
b.pPackets = nil
b.bound = true
b.logger.V(1).Info("NewBuffer", "MaxBitRate", o.MaxBitRate)
}
// Write adds a RTP Packet, out of order, new packet may be arrived later
func (b *Buffer) Write(pkt []byte) (n int, err error) {
b.Lock()
defer b.Unlock()
if b.closed.get() {
err = io.EOF
return
}
if !b.bound {
packet := make([]byte, len(pkt))
copy(packet, pkt)
b.pPackets = append(b.pPackets, pendingPackets{
packet: packet,
arrivalTime: time.Now().UnixNano(),
})
return
}
b.calc(pkt, time.Now().UnixNano())
return
}
func (b *Buffer) Read(buff []byte) (n int, err error) {
for {
if b.closed.get() {
err = io.EOF
return
}
b.Lock()
if b.pPackets != nil && len(b.pPackets) > b.lastPacketRead {
if len(buff) < len(b.pPackets[b.lastPacketRead].packet) {
err = errBufferTooSmall
b.Unlock()
return
}
n = len(b.pPackets[b.lastPacketRead].packet)
copy(buff, b.pPackets[b.lastPacketRead].packet)
b.lastPacketRead++
b.Unlock()
return
}
b.Unlock()
time.Sleep(25 * time.Millisecond)
}
}
func (b *Buffer) ReadExtended() (*ExtPacket, error) {
for {
if b.closed.get() {
return nil, io.EOF
}
b.Lock()
if b.extPackets.Len() > 0 {
extPkt := b.extPackets.PopFront().(*ExtPacket)
b.Unlock()
return extPkt, nil
}
b.Unlock()
time.Sleep(10 * time.Millisecond)
}
}
func (b *Buffer) Close() error {
b.Lock()
defer b.Unlock()
b.closeOnce.Do(func() {
if b.bucket != nil && b.codecType == webrtc.RTPCodecTypeVideo {
b.videoPool.Put(b.bucket.src)
}
if b.bucket != nil && b.codecType == webrtc.RTPCodecTypeAudio {
b.audioPool.Put(b.bucket.src)
}
b.closed.set(true)
b.onClose()
})
return nil
}
func (b *Buffer) OnClose(fn func()) {
b.onClose = fn
}
func (b *Buffer) calc(pkt []byte, arrivalTime int64) {
sn := binary.BigEndian.Uint16(pkt[2:4])
var headPkt bool
if b.stats.PacketCount == 0 {
b.baseSN = sn
b.lastReport = arrivalTime
b.seqHdlr.UpdateMaxSeq(uint32(sn))
headPkt = true
} else {
extSN, isNewer := b.seqHdlr.Unwrap(sn)
if b.nack {
if isNewer {
for i := b.seqHdlr.MaxSeqNo() + 1; i < extSN; i++ {
b.nacker.push(i)
}
} else {
b.nacker.remove(extSN)
}
}
if isNewer {
b.seqHdlr.UpdateMaxSeq(extSN)
}
headPkt = isNewer
}
var p rtp.Packet
pb, err := b.bucket.AddPacket(pkt, sn, headPkt)
if err != nil {
if err == errRTXPacket {
return
}
return
}
if err = p.Unmarshal(pb); err != nil {
return
}
// submit to TWCC even if it is a padding only packet. Clients use padding only packets as probes
// for bandwidth estimation
if b.twcc {
if ext := p.GetExtension(b.twccExt); ext != nil && len(ext) > 1 {
b.feedbackTWCC(binary.BigEndian.Uint16(ext[0:2]), arrivalTime, p.Marker)
}
}
b.stats.TotalByte += uint64(len(pkt))
b.stats.PacketCount++
ep := ExtPacket{
Head: headPkt,
Cycle: b.seqHdlr.Cycles(),
Packet: p,
Arrival: arrivalTime,
}
if len(p.Payload) == 0 {
// padding only packet, nothing else to do
b.extPackets.PushBack(&ep)
return
}
temporalLayer := int32(0)
switch b.mime {
case "video/vp8":
vp8Packet := VP8{}
if err := vp8Packet.Unmarshal(p.Payload); err != nil {
return
}
ep.Payload = vp8Packet
ep.KeyFrame = vp8Packet.IsKeyFrame
temporalLayer = int32(vp8Packet.TID)
case "video/h264":
ep.KeyFrame = isH264Keyframe(p.Payload)
}
if b.minPacketProbe < 25 {
if sn < b.baseSN {
b.baseSN = sn
}
if b.mime == "video/vp8" {
pld := ep.Payload.(VP8)
mtl := atomic.LoadInt32(&b.maxTemporalLayer)
if mtl < int32(pld.TID) {
atomic.StoreInt32(&b.maxTemporalLayer, int32(pld.TID))
}
}
b.minPacketProbe++
}
b.extPackets.PushBack(&ep)
// if first time update or the timestamp is later (factoring timestamp wrap around)
latestTimestamp := atomic.LoadUint32(&b.latestTimestamp)
latestTimestampTimeInNanosSinceEpoch := atomic.LoadInt64(&b.latestTimestampTime)
if (latestTimestampTimeInNanosSinceEpoch == 0) || IsLaterTimestamp(p.Timestamp, latestTimestamp) {
atomic.StoreUint32(&b.latestTimestamp, p.Timestamp)
atomic.StoreInt64(&b.latestTimestampTime, arrivalTime)
}
arrival := uint32(arrivalTime / 1e6 * int64(b.clockRate/1e3))
transit := arrival - p.Timestamp
if b.lastTransit != 0 {
d := int32(transit - b.lastTransit)
if d < 0 {
d = -d
}
b.stats.Jitter += (float64(d) - b.stats.Jitter) / 16
}
b.lastTransit = transit
if b.audioLevel {
if e := p.GetExtension(b.audioExt); e != nil && b.onAudioLevel != nil {
ext := rtp.AudioLevelExtension{}
if err := ext.Unmarshal(e); err == nil {
b.onAudioLevel(ext.Level)
}
}
}
if b.nacker != nil {
if r := b.buildNACKPacket(); r != nil {
b.feedbackCB(r)
}
}
b.bitrateHelper[temporalLayer] += uint64(len(pkt))
diff := arrivalTime - b.lastReport
if diff >= reportDelta {
// LK-TODO-START
// As this happens in the data path, if there are no packets received
// in an interval, the bitrate is stuck with the old value. GetBitrate()
// method in sfu.Receiver uses the availableLayers set by stream
// tracker to report 0 bitrate if a layer is not available. But, stream
// tracker is not run for the lowest layer. So, if the lowest layer stops,
// stale bitrate will be reported. The simplest thing might be to run the
// stream tracker on all layers to address this. Another option to look at
// is some monitoring loop running at low frequency and reporting bitrate.
// LK-TODO-END
bitrates, ok := b.bitrate.Load().([]uint64)
if !ok {
bitrates = make([]uint64, len(b.bitrateHelper))
}
for i := 0; i < len(b.bitrateHelper); i++ {
br := (8 * b.bitrateHelper[i] * uint64(reportDelta)) / uint64(diff)
bitrates[i] = br
b.bitrateHelper[i] = 0
}
b.bitrate.Store(bitrates)
b.feedbackCB(b.getRTCP())
b.lastReport = arrivalTime
}
}
func (b *Buffer) buildNACKPacket() []rtcp.Packet {
if nacks, askKeyframe := b.nacker.pairs(b.seqHdlr.MaxSeqNo()); (nacks != nil && len(nacks) > 0) || askKeyframe {
var pkts []rtcp.Packet
if len(nacks) > 0 {
pkts = []rtcp.Packet{&rtcp.TransportLayerNack{
MediaSSRC: b.mediaSSRC,
Nacks: nacks,
}}
}
if askKeyframe {
pkts = append(pkts, &rtcp.PictureLossIndication{
MediaSSRC: b.mediaSSRC,
})
}
return pkts
}
return nil
}
func (b *Buffer) buildREMBPacket() *rtcp.ReceiverEstimatedMaximumBitrate {
br := b.Bitrate()
if b.stats.LostRate < 0.02 {
br = uint64(float64(br)*1.09) + 2000
}
if b.stats.LostRate > .1 {
br = uint64(float64(br) * float64(1-0.5*b.stats.LostRate))
}
if br > b.maxBitrate {
br = b.maxBitrate
}
if br < 100000 {
br = 100000
}
b.stats.TotalByte = 0
return &rtcp.ReceiverEstimatedMaximumBitrate{
Bitrate: float32(br),
SSRCs: []uint32{b.mediaSSRC},
}
}
func (b *Buffer) buildReceptionReport() rtcp.ReceptionReport {
extMaxSeq := b.seqHdlr.MaxSeqNo()
expected := extMaxSeq - uint32(b.baseSN) + 1
lost := uint32(0)
if b.stats.PacketCount < expected && b.stats.PacketCount != 0 {
lost = expected - b.stats.PacketCount
}
expectedInterval := expected - b.stats.LastExpected
b.stats.LastExpected = expected
receivedInterval := b.stats.PacketCount - b.stats.LastReceived
b.stats.LastReceived = b.stats.PacketCount
lostInterval := expectedInterval - receivedInterval
b.stats.LostRate = float32(lostInterval) / float32(expectedInterval)
var fracLost uint8
if expectedInterval != 0 && lostInterval > 0 {
fracLost = uint8((lostInterval << 8) / expectedInterval)
}
if b.lastFractionLostToReport > fracLost {
// If fractionlost from subscriber is bigger than sfu received, use it.
fracLost = b.lastFractionLostToReport
}
var dlsr uint32
if b.lastSRRecv != 0 {
delayMS := uint32((time.Now().UnixNano() - b.lastSRRecv) / 1e6)
dlsr = (delayMS / 1e3) << 16
dlsr |= (delayMS % 1e3) * 65536 / 1000
}
rr := rtcp.ReceptionReport{
SSRC: b.mediaSSRC,
FractionLost: fracLost,
TotalLost: lost,
LastSequenceNumber: extMaxSeq,
Jitter: uint32(b.stats.Jitter),
LastSenderReport: uint32(b.lastSRNTPTime >> 16),
Delay: dlsr,
}
return rr
}
func (b *Buffer) SetSenderReportData(rtpTime uint32, ntpTime uint64) {
b.Lock()
b.lastSRRTPTime = rtpTime
b.lastSRNTPTime = ntpTime
b.lastSRRecv = time.Now().UnixNano()
b.Unlock()
}
func (b *Buffer) SetLastFractionLostReport(lost uint8) {
b.lastFractionLostToReport = lost
}
func (b *Buffer) getRTCP() []rtcp.Packet {
var pkts []rtcp.Packet
pkts = append(pkts, &rtcp.ReceiverReport{
Reports: []rtcp.ReceptionReport{b.buildReceptionReport()},
})
if b.remb && !b.twcc {
pkts = append(pkts, b.buildREMBPacket())
}
return pkts
}
func (b *Buffer) GetPacket(buff []byte, sn uint16) (int, error) {
b.Lock()
defer b.Unlock()
if b.closed.get() {
return 0, io.EOF
}
return b.bucket.GetPacket(buff, sn)
}
// Bitrate returns the current publisher stream bitrate.
func (b *Buffer) Bitrate() uint64 {
bitrates, ok := b.bitrate.Load().([]uint64)
bitrate := uint64(0)
if ok {
for _, b := range bitrates {
bitrate += b
}
}
return bitrate
}
// BitrateTemporal returns the current publisher stream bitrate temporal layer wise.
func (b *Buffer) BitrateTemporal() []uint64 {
bitrates, ok := b.bitrate.Load().([]uint64)
if !ok {
return make([]uint64, len(b.bitrateHelper))
}
// copy and return
brs := make([]uint64, len(bitrates))
copy(brs, bitrates)
return brs
}
// BitrateTemporalCumulative returns the current publisher stream bitrate temporal layer accumulated with lower temporal layers.
func (b *Buffer) BitrateTemporalCumulative() []uint64 {
bitrates, ok := b.bitrate.Load().([]uint64)
if !ok {
return make([]uint64, len(b.bitrateHelper))
}
// copy and process
brs := make([]uint64, len(bitrates))
copy(brs, bitrates)
for i := len(brs) - 1; i >= 1; i-- {
if brs[i] != 0 {
for j := i - 1; j >= 0; j-- {
brs[i] += brs[j]
}
}
}
return brs
}
func (b *Buffer) MaxTemporalLayer() int32 {
return atomic.LoadInt32(&b.maxTemporalLayer)
}
func (b *Buffer) OnTransportWideCC(fn func(sn uint16, timeNS int64, marker bool)) {
b.feedbackTWCC = fn
}
func (b *Buffer) OnFeedback(fn func(fb []rtcp.Packet)) {
b.feedbackCB = fn
}
func (b *Buffer) OnAudioLevel(fn func(level uint8)) {
b.onAudioLevel = fn
}
// 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
}
// GetSenderReportData returns the rtp, ntp and nanos of the last sender report
func (b *Buffer) GetSenderReportData() (rtpTime uint32, ntpTime uint64, lastReceivedTimeInNanosSinceEpoch int64) {
rtpTime = atomic.LoadUint32(&b.lastSRRTPTime)
ntpTime = atomic.LoadUint64(&b.lastSRNTPTime)
lastReceivedTimeInNanosSinceEpoch = atomic.LoadInt64(&b.lastSRRecv)
return rtpTime, ntpTime, lastReceivedTimeInNanosSinceEpoch
}
// GetStats returns the raw statistics about a particular buffer state
func (b *Buffer) GetStats() (stats Stats) {
b.Lock()
stats = b.stats
b.Unlock()
return
}
// GetLatestTimestamp returns the latest RTP timestamp factoring in potential RTP timestamp wrap-around
func (b *Buffer) GetLatestTimestamp() (latestTimestamp uint32, latestTimestampTimeInNanosSinceEpoch int64) {
latestTimestamp = atomic.LoadUint32(&b.latestTimestamp)
latestTimestampTimeInNanosSinceEpoch = atomic.LoadInt64(&b.latestTimestampTime)
return latestTimestamp, latestTimestampTimeInNanosSinceEpoch
}
// IsTimestampWrapAround returns true if wrap around happens from timestamp1 to timestamp2
func IsTimestampWrapAround(timestamp1 uint32, timestamp2 uint32) bool {
return timestamp2 < timestamp1 && timestamp1 > 0xf0000000 && timestamp2 < 0x0fffffff
}
// IsLaterTimestamp returns true if timestamp1 is later in time than timestamp2 factoring in timestamp wrap-around
func IsLaterTimestamp(timestamp1 uint32, timestamp2 uint32) bool {
if timestamp1 > timestamp2 {
if IsTimestampWrapAround(timestamp1, timestamp2) {
return false
}
return true
}
if IsTimestampWrapAround(timestamp2, timestamp1) {
return true
}
return false
}
func isNewerUint16(val1, val2 uint16) bool {
return val1 != val2 && val1-val2 < 0x8000
}
type SeqWrapHandler struct {
maxSeqNo uint32
}
func (s *SeqWrapHandler) Cycles() uint32 {
return s.maxSeqNo & 0xffff0000
}
func (s *SeqWrapHandler) MaxSeqNo() uint32 {
return s.maxSeqNo
}
// unwrap seq and update the maxSeqNo. return unwraped value, and whether seq is newer
func (s *SeqWrapHandler) Unwrap(seq uint16) (uint32, bool) {
maxSeqNo := uint16(s.maxSeqNo)
delta := int32(seq) - int32(maxSeqNo)
newer := isNewerUint16(seq, maxSeqNo)
if newer {
if delta < 0 {
// seq is newer, but less than maxSeqNo, wrap around
delta += 0x10000
}
} else {
// older value
if delta > 0 && (int32(s.maxSeqNo)+delta-0x10000) >= 0 {
// wrap backwards, should not less than 0 in this case:
// at start time, received seq 1, set s.maxSeqNo =1 ,
// then a out of order seq 65534 coming, we can't unwrap
// the seq to -2
delta -= 0x10000
}
}
unwrapped := uint32(int32(s.maxSeqNo) + delta)
return unwrapped, newer
}
func (s *SeqWrapHandler) UpdateMaxSeq(extSeq uint32) {
s.maxSeqNo = extSeq
}
+363
View File
@@ -0,0 +1,363 @@
package buffer
import (
"sync"
"testing"
"time"
"github.com/livekit/livekit-server/pkg/sfu/logger"
"github.com/pion/rtcp"
"github.com/pion/rtp"
"github.com/pion/webrtc/v3"
"github.com/stretchr/testify/assert"
)
func CreateTestPacket(pktStamp *SequenceNumberAndTimeStamp) *rtp.Packet {
if pktStamp == nil {
return &rtp.Packet{
Header: rtp.Header{},
Payload: []byte{1, 2, 3},
}
}
return &rtp.Packet{
Header: rtp.Header{
SequenceNumber: pktStamp.SequenceNumber,
Timestamp: pktStamp.Timestamp,
},
Payload: []byte{1, 2, 3},
}
}
type SequenceNumberAndTimeStamp struct {
SequenceNumber uint16
Timestamp uint32
}
func CreateTestListPackets(snsAndTSs []SequenceNumberAndTimeStamp) (packetList []*rtp.Packet) {
for _, item := range snsAndTSs {
item := item
packetList = append(packetList, CreateTestPacket(&item))
}
return packetList
}
func TestNack(t *testing.T) {
pool := &sync.Pool{
New: func() interface{} {
b := make([]byte, 1500)
return &b
},
}
logger.SetGlobalOptions(logger.GlobalConfig{V: 1}) // 2 - TRACE
logger := logger.New()
t.Run("nack normal", func(t *testing.T) {
buff := NewBuffer(123, pool, pool, logger)
buff.codecType = webrtc.RTPCodecTypeVideo
assert.NotNil(t, buff)
var wg sync.WaitGroup
// 3 nacks 1 Pli
wg.Add(4)
buff.OnFeedback(func(fb []rtcp.Packet) {
for _, pkt := range fb {
switch p := pkt.(type) {
case *rtcp.TransportLayerNack:
if p.Nacks[0].PacketList()[0] == 1 && p.MediaSSRC == 123 {
wg.Done()
}
case *rtcp.PictureLossIndication:
if p.MediaSSRC == 123 {
wg.Done()
}
}
}
})
buff.Bind(webrtc.RTPParameters{
HeaderExtensions: nil,
Codecs: []webrtc.RTPCodecParameters{
{
RTPCodecCapability: webrtc.RTPCodecCapability{
MimeType: "video/vp8",
ClockRate: 90000,
RTCPFeedback: []webrtc.RTCPFeedback{{
Type: "nack",
}},
},
PayloadType: 96,
},
},
}, Options{})
for i := 0; i < 15; i++ {
if i == 1 {
continue
}
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i), Timestamp: uint32(i)},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
assert.NoError(t, err)
_, err = buff.Write(b)
assert.NoError(t, err)
}
wg.Wait()
})
t.Run("nack with seq wrap", func(t *testing.T) {
buff := NewBuffer(123, pool, pool, logger)
buff.codecType = webrtc.RTPCodecTypeVideo
assert.NotNil(t, buff)
var wg sync.WaitGroup
expects := map[uint16]int{
65534: 0,
65535: 0,
0: 0,
1: 0,
}
wg.Add(3 * len(expects)) // retry 3 times
buff.OnFeedback(func(fb []rtcp.Packet) {
for _, pkt := range fb {
switch p := pkt.(type) {
case *rtcp.TransportLayerNack:
if p.MediaSSRC == 123 {
for _, v := range p.Nacks {
v.Range(func(seq uint16) bool {
if _, ok := expects[seq]; ok {
wg.Done()
} else {
assert.Fail(t, "unexpected nack seq ", seq)
}
return true
})
}
}
case *rtcp.PictureLossIndication:
if p.MediaSSRC == 123 {
// wg.Done()
}
}
}
})
buff.Bind(webrtc.RTPParameters{
HeaderExtensions: nil,
Codecs: []webrtc.RTPCodecParameters{
{
RTPCodecCapability: webrtc.RTPCodecCapability{
MimeType: "video/vp8",
ClockRate: 90000,
RTCPFeedback: []webrtc.RTCPFeedback{{
Type: "nack",
}},
},
PayloadType: 96,
},
},
}, Options{})
for i := 0; i < 15; i++ {
if i > 0 && i < 5 {
continue
}
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i + 65533), Timestamp: uint32(i)},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
assert.NoError(t, err)
_, err = buff.Write(b)
assert.NoError(t, err)
}
wg.Wait()
})
}
func TestNewBuffer(t *testing.T) {
type args struct {
options Options
ssrc uint32
}
tests := []struct {
name string
args args
}{
{
name: "Must not be nil and add packets in sequence",
args: args{
options: Options{
MaxBitRate: 1e6,
},
},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
var TestPackets = []*rtp.Packet{
{
Header: rtp.Header{
SequenceNumber: 65533,
},
},
{
Header: rtp.Header{
SequenceNumber: 65534,
},
},
{
Header: rtp.Header{
SequenceNumber: 2,
},
},
{
Header: rtp.Header{
SequenceNumber: 65535,
},
},
}
pool := &sync.Pool{
New: func() interface{} {
b := make([]byte, 1500)
return &b
},
}
logger.SetGlobalOptions(logger.GlobalConfig{V: 2}) // 2 - TRACE
logger := logger.New()
buff := NewBuffer(123, pool, pool, logger)
buff.codecType = webrtc.RTPCodecTypeVideo
assert.NotNil(t, buff)
assert.NotNil(t, TestPackets)
buff.OnFeedback(func(_ []rtcp.Packet) {
})
buff.Bind(webrtc.RTPParameters{
HeaderExtensions: nil,
Codecs: []webrtc.RTPCodecParameters{{
RTPCodecCapability: webrtc.RTPCodecCapability{
MimeType: "video/vp8",
ClockRate: 9600,
RTCPFeedback: nil,
},
PayloadType: 0,
}},
}, Options{})
for _, p := range TestPackets {
buf, _ := p.Marshal()
buff.Write(buf)
}
// assert.Equal(t, 6, buff.PacketQueue.size)
assert.Equal(t, uint32(1<<16), buff.seqHdlr.Cycles())
assert.Equal(t, uint16(2), uint16(buff.seqHdlr.MaxSeqNo()))
})
}
}
func TestFractionLostReport(t *testing.T) {
pool := &sync.Pool{
New: func() interface{} {
b := make([]byte, 1500)
return &b
},
}
logger.SetGlobalOptions(logger.GlobalConfig{V: 1}) // 2 - TRACE
buff := NewBuffer(123, pool, pool, logger.New())
buff.codecType = webrtc.RTPCodecTypeVideo
assert.NotNil(t, buff)
var wg sync.WaitGroup
wg.Add(1)
buff.SetLastFractionLostReport(55)
buff.OnFeedback(func(fb []rtcp.Packet) {
for _, pkt := range fb {
switch p := pkt.(type) {
case *rtcp.ReceiverReport:
for _, v := range p.Reports {
assert.EqualValues(t, 55, v.FractionLost)
}
wg.Done()
}
}
})
buff.Bind(webrtc.RTPParameters{
HeaderExtensions: nil,
Codecs: []webrtc.RTPCodecParameters{
{
RTPCodecCapability: webrtc.RTPCodecCapability{
MimeType: "audio/opus",
ClockRate: 48000,
},
PayloadType: 96,
},
},
}, Options{})
for i := 0; i < 15; i++ {
pkt := rtp.Packet{
Header: rtp.Header{SequenceNumber: uint16(i), Timestamp: uint32(i)},
Payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1},
}
b, err := pkt.Marshal()
assert.NoError(t, err)
if i == 1 {
time.Sleep(1 * time.Second)
}
_, err = buff.Write(b)
assert.NoError(t, err)
}
wg.Wait()
}
func TestSeqWrapHandler(t *testing.T) {
s := SeqWrapHandler{}
s.UpdateMaxSeq(1)
assert.Equal(t, uint32(1), s.MaxSeqNo())
type caseInfo struct {
seqs []uint32 //{seq1, seq2, unwrap of seq2}
newer bool // seq2 is newer than seq1
}
// test normal case, name -> {seq1, seq2, unwrap of seq2}
cases := map[string]caseInfo{
"no wrap": {[]uint32{1, 4, 4}, true},
"no wrap backward": {[]uint32{4, 1, 1}, false},
"wrap around forward to zero": {[]uint32{65534, 0, 65536}, true},
"wrap around forward": {[]uint32{65534, 10, 65546}, true},
"wrap around forward 2": {[]uint32{65535 + 65536*2, 1, 1 + 65536*3}, true},
"wrap around backward ": {[]uint32{5, 65534, 65534}, false},
"wrap around backward less than zero": {[]uint32{5, 65534, 65534}, false},
}
for k, v := range cases {
t.Run(k, func(t *testing.T) {
s := SeqWrapHandler{}
s.UpdateMaxSeq(v.seqs[0])
extsn, newer := s.Unwrap(uint16(v.seqs[1]))
assert.Equal(t, v.newer, newer)
assert.Equal(t, v.seqs[2], extsn)
})
}
}
func TestIsTimestampWrap(t *testing.T) {
type caseInfo struct {
name string
ts1 uint32
ts2 uint32
later bool
}
cases := []caseInfo{
{"normal case 1 timestamp later ", 2, 1, true},
{"normal case 2 timestamp later", 0x1c000000, 0x10000000, true},
{"wrap case timestamp later", 0xffff, 0xfc000000, true},
{"wrap case timestamp early", 0xfc000000, 0xffff, false},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
assert.Equal(t, c.later, IsLaterTimestamp(c.ts1, c.ts2))
})
}
}
+10
View File
@@ -0,0 +1,10 @@
package buffer
import "errors"
var (
errPacketNotFound = errors.New("packet not found in cache")
errBufferTooSmall = errors.New("buffer too small")
errPacketTooOld = errors.New("received packet too old")
errRTXPacket = errors.New("packet already received")
)
+97
View File
@@ -0,0 +1,97 @@
package buffer
import (
"io"
"sync"
"github.com/go-logr/logr"
"github.com/pion/transport/packetio"
)
type Factory struct {
sync.RWMutex
videoPool *sync.Pool
audioPool *sync.Pool
rtpBuffers map[uint32]*Buffer
rtcpReaders map[uint32]*RTCPReader
logger logr.Logger
}
func NewBufferFactory(trackingPackets int, logger logr.Logger) *Factory {
// Enable package wide logging for non-method functions.
// If logger is empty - use default Logger.
// Logger is a public variable in buffer package.
if logger == (logr.Logger{}) {
logger = Logger
} else {
Logger = logger
}
return &Factory{
videoPool: &sync.Pool{
New: func() interface{} {
b := make([]byte, trackingPackets*maxPktSize)
return &b
},
},
audioPool: &sync.Pool{
New: func() interface{} {
b := make([]byte, maxPktSize*25)
return &b
},
},
rtpBuffers: make(map[uint32]*Buffer),
rtcpReaders: make(map[uint32]*RTCPReader),
logger: logger,
}
}
func (f *Factory) GetOrNew(packetType packetio.BufferPacketType, ssrc uint32) io.ReadWriteCloser {
f.Lock()
defer f.Unlock()
switch packetType {
case packetio.RTCPBufferPacket:
if reader, ok := f.rtcpReaders[ssrc]; ok {
return reader
}
reader := NewRTCPReader(ssrc)
f.rtcpReaders[ssrc] = reader
reader.OnClose(func() {
f.Lock()
delete(f.rtcpReaders, ssrc)
f.Unlock()
})
return reader
case packetio.RTPBufferPacket:
if reader, ok := f.rtpBuffers[ssrc]; ok {
return reader
}
buffer := NewBuffer(ssrc, f.videoPool, f.audioPool, f.logger)
f.rtpBuffers[ssrc] = buffer
buffer.OnClose(func() {
f.Lock()
delete(f.rtpBuffers, ssrc)
f.Unlock()
})
return buffer
}
return nil
}
func (f *Factory) GetBufferPair(ssrc uint32) (*Buffer, *RTCPReader) {
f.RLock()
defer f.RUnlock()
return f.rtpBuffers[ssrc], f.rtcpReaders[ssrc]
}
func (f *Factory) GetBuffer(ssrc uint32) *Buffer {
f.RLock()
defer f.RUnlock()
return f.rtpBuffers[ssrc]
}
func (f *Factory) GetRTCPReader(ssrc uint32) *RTCPReader {
f.RLock()
defer f.RUnlock()
return f.rtcpReaders[ssrc]
}
+297
View File
@@ -0,0 +1,297 @@
package buffer
import (
"encoding/binary"
"errors"
"sync/atomic"
)
var (
errShortPacket = errors.New("packet is not large enough")
errNilPacket = errors.New("invalid nil packet")
errInvalidPacket = errors.New("invalid packet")
)
type atomicBool int32
func (a *atomicBool) set(value bool) {
var i int32
if value {
i = 1
}
atomic.StoreInt32((*int32)(a), i)
}
func (a *atomicBool) get() bool {
return atomic.LoadInt32((*int32)(a)) != 0
}
// VP8 is a helper to get temporal data from VP8 packet header
/*
VP8 Payload Descriptor
0 1 2 3 4 5 6 7 0 1 2 3 4 5 6 7
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
|X|R|N|S|R| PID | (REQUIRED) |X|R|N|S|R| PID | (REQUIRED)
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
X: |I|L|T|K| RSV | (OPTIONAL) X: |I|L|T|K| RSV | (OPTIONAL)
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
I: |M| PictureID | (OPTIONAL) I: |M| PictureID | (OPTIONAL)
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
L: | TL0PICIDX | (OPTIONAL) | PictureID |
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
T/K:|TID|Y| KEYIDX | (OPTIONAL) L: | TL0PICIDX | (OPTIONAL)
+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+
T/K:|TID|Y| KEYIDX | (OPTIONAL)
+-+-+-+-+-+-+-+-+
*/
type VP8 struct {
TemporalSupported bool // LK-TODO: CLEANUP-REMOVE
FirstByte byte
PictureIDPresent int
PictureID uint16 /* 8 or 16 bits, picture ID */
PicIDIdx int // LK-TODO: CLEANUP-REMOVE
MBit bool
TL0PICIDXPresent int
TL0PICIDX uint8 /* 8 bits temporal level zero index */
TlzIdx int // LK-TODO: CLEANUP-REMOVE
// Optional Header If either of the T or K bits are set to 1,
// the TID/Y/KEYIDX extension field MUST be present.
TIDPresent int
TID uint8 /* 2 bits temporal layer idx */
Y uint8
KEYIDXPresent int
KEYIDX uint8 /* 5 bits of key frame idx */
HeaderSize int
// IsKeyFrame is a helper to detect if current packet is a keyframe
IsKeyFrame bool
}
// Unmarshal parses the passed byte slice and stores the result in the VP8 this method is called upon
func (p *VP8) Unmarshal(payload []byte) error {
if payload == nil {
return errNilPacket
}
payloadLen := len(payload)
if payloadLen < 1 {
return errShortPacket
}
idx := 0
p.FirstByte = payload[idx]
S := payload[idx]&0x10 > 0
// Check for extended bit control
if payload[idx]&0x80 > 0 {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
I := payload[idx]&0x80 > 0
L := payload[idx]&0x40 > 0
T := payload[idx]&0x20 > 0
K := payload[idx]&0x10 > 0
if L && !T {
return errInvalidPacket
}
// Check if T is present, if not, no temporal layer is available
p.TemporalSupported = payload[idx]&0x20 > 0
// Check for PictureID
if I {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
p.PicIDIdx = idx
p.PictureIDPresent = 1
pid := payload[idx] & 0x7f
// Check if m is 1, then Picture ID is 15 bits
if payload[idx]&0x80 > 0 {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
p.MBit = true
p.PictureID = binary.BigEndian.Uint16([]byte{pid, payload[idx]})
} else {
p.PictureID = uint16(pid)
}
}
// Check if TL0PICIDX is present
if L {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
p.TlzIdx = idx
p.TL0PICIDXPresent = 1
if int(idx) >= payloadLen {
return errShortPacket
}
p.TL0PICIDX = payload[idx]
}
if T || K {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
if T {
p.TIDPresent = 1
p.TID = (payload[idx] & 0xc0) >> 6
p.Y = (payload[idx] & 0x20) >> 5
}
if K {
p.KEYIDXPresent = 1
p.KEYIDX = payload[idx] & 0x1f
}
}
if idx >= payloadLen {
return errShortPacket
}
idx++
if payloadLen < idx+1 {
return errShortPacket
}
// Check is packet is a keyframe by looking at P bit in vp8 payload
p.IsKeyFrame = payload[idx]&0x01 == 0 && S
} else {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
// Check is packet is a keyframe by looking at P bit in vp8 payload
p.IsKeyFrame = payload[idx]&0x01 == 0 && S
}
p.HeaderSize = idx
return nil
}
func (v *VP8) MarshalTo(buf []byte) error {
if len(buf) < v.HeaderSize {
return errShortPacket
}
idx := 0
buf[idx] = v.FirstByte
if (v.PictureIDPresent + v.TL0PICIDXPresent + v.TIDPresent + v.KEYIDXPresent) != 0 {
buf[idx] |= 0x80 // X bit
idx++
buf[idx] = byte(v.PictureIDPresent<<7) | byte(v.TL0PICIDXPresent<<6) | byte(v.TIDPresent<<5) | byte(v.KEYIDXPresent<<4)
idx++
if v.PictureIDPresent == 1 {
if v.MBit {
buf[idx] = 0x80 | byte((v.PictureID>>8)&0x7f)
buf[idx+1] = byte(v.PictureID & 0xff)
idx += 2
} else {
buf[idx] = byte(v.PictureID)
idx++
}
}
if v.TL0PICIDXPresent == 1 {
buf[idx] = byte(v.TL0PICIDX)
idx++
}
if v.TIDPresent == 1 || v.KEYIDXPresent == 1 {
buf[idx] = 0
if v.TIDPresent == 1 {
buf[idx] = byte(v.TID<<6) | byte(v.Y<<5)
}
if v.KEYIDXPresent == 1 {
buf[idx] |= byte(v.KEYIDX & 0x1f)
}
idx++
}
} else {
buf[idx] &^= 0x80 // X bit
idx++
}
return nil
}
func VP8PictureIdSizeDiff(mBit1 bool, mBit2 bool) int {
if mBit1 == mBit2 {
return 0
}
if mBit1 {
return 1
}
return -1
}
// isH264Keyframe detects if h264 payload is a keyframe
// this code was taken from https://github.com/jech/galene/blob/codecs/rtpconn/rtpreader.go#L45
// all credits belongs to Juliusz Chroboczek @jech and the awesome Galene SFU
func isH264Keyframe(payload []byte) bool {
if len(payload) < 1 {
return false
}
nalu := payload[0] & 0x1F
if nalu == 0 {
// reserved
return false
} else if nalu <= 23 {
// simple NALU
return nalu == 5
} else if nalu == 24 || nalu == 25 || nalu == 26 || nalu == 27 {
// STAP-A, STAP-B, MTAP16 or MTAP24
i := 1
if nalu == 25 || nalu == 26 || nalu == 27 {
// skip DON
i += 2
}
for i < len(payload) {
if i+2 > len(payload) {
return false
}
length := uint16(payload[i])<<8 |
uint16(payload[i+1])
i += 2
if i+int(length) > len(payload) {
return false
}
offset := 0
if nalu == 26 {
offset = 3
} else if nalu == 27 {
offset = 4
}
if offset >= int(length) {
return false
}
n := payload[i+offset] & 0x1F
if n == 7 {
return true
} else if n >= 24 {
// is this legal?
Logger.V(0).Info("Non-simple NALU within a STAP")
}
i += int(length)
}
if i == len(payload) {
return false
}
return false
} else if nalu == 28 || nalu == 29 {
// FU-A or FU-B
if len(payload) < 2 {
return false
}
if (payload[1] & 0x80) == 0 {
// not a starting fragment
return false
}
return payload[1]&0x1F == 7
}
return false
}
+94
View File
@@ -0,0 +1,94 @@
package buffer
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestVP8Helper_Unmarshal(t *testing.T) {
type args struct {
payload []byte
}
tests := []struct {
name string
args args
wantErr bool
checkTemporal bool
temporalSupport bool
checkKeyFrame bool
keyFrame bool
checkPictureID bool
pictureID uint16
checkTlzIdx bool
tlzIdx uint8
checkTempID bool
temporalID uint8
}{
{
name: "Empty or nil payload must return error",
args: args{payload: []byte{}},
wantErr: true,
},
{
name: "Temporal must be supported by setting T bit to 1",
args: args{payload: []byte{0xff, 0x20, 0x1, 0x2, 0x3, 0x4}},
checkTemporal: true,
temporalSupport: true,
},
{
name: "Picture must be ID 7 bits by setting M bit to 0 and present by I bit set to 1",
args: args{payload: []byte{0xff, 0xff, 0x11, 0x2, 0x3, 0x4}},
checkPictureID: true,
pictureID: 17,
},
{
name: "Picture ID must be 15 bits by setting M bit to 1 and present by I bit set to 1",
args: args{payload: []byte{0xff, 0xff, 0x92, 0x67, 0x3, 0x4, 0x5}},
checkPictureID: true,
pictureID: 4711,
},
{
name: "Temporal level zero index must be present if L set to 1",
args: args{payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x4, 0x5}},
checkTlzIdx: true,
tlzIdx: 180,
},
{
name: "Temporal index must be present and used if T bit set to 1",
args: args{payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x5, 0x6}},
checkTempID: true,
temporalID: 2,
},
{
name: "Check if packet is a keyframe by looking at P bit set to 0",
args: args{payload: []byte{0xff, 0xff, 0xff, 0xfd, 0xb4, 0x9f, 0x94, 0x1}},
checkKeyFrame: true,
keyFrame: true,
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
p := &VP8{}
if err := p.Unmarshal(tt.args.payload); (err != nil) != tt.wantErr {
t.Errorf("Unmarshal() error = %v, wantErr %v", err, tt.wantErr)
}
if tt.checkTemporal {
assert.Equal(t, tt.temporalSupport, p.TemporalSupported)
}
if tt.checkKeyFrame {
assert.Equal(t, tt.keyFrame, p.IsKeyFrame)
}
if tt.checkPictureID {
assert.Equal(t, tt.pictureID, p.PictureID)
}
if tt.checkTlzIdx {
assert.Equal(t, tt.tlzIdx, p.TL0PICIDX)
}
if tt.checkTempID {
assert.Equal(t, tt.temporalID, p.TID)
}
})
}
}
+106
View File
@@ -0,0 +1,106 @@
package buffer
import (
"sort"
"github.com/pion/rtcp"
)
const maxNackTimes = 3 // Max number of times a packet will be NACKed
const maxNackCache = 100 // Max NACK sn the sfu will keep reference
type nack struct {
sn uint32
nacked uint8
}
type nackQueue struct {
nacks []nack
kfSN uint32
}
func newNACKQueue() *nackQueue {
return &nackQueue{
nacks: make([]nack, 0, maxNackCache+1),
}
}
func (n *nackQueue) remove(extSN uint32) {
i := sort.Search(len(n.nacks), func(i int) bool { return n.nacks[i].sn >= extSN })
if i >= len(n.nacks) || n.nacks[i].sn != extSN {
return
}
copy(n.nacks[i:], n.nacks[i+1:])
n.nacks = n.nacks[:len(n.nacks)-1]
}
func (n *nackQueue) push(extSN uint32) {
i := sort.Search(len(n.nacks), func(i int) bool { return n.nacks[i].sn >= extSN })
if i < len(n.nacks) && n.nacks[i].sn == extSN {
return
}
nck := nack{
sn: extSN,
nacked: 0,
}
if i == len(n.nacks) {
n.nacks = append(n.nacks, nck)
} else {
n.nacks = append(n.nacks[:i+1], n.nacks[i:]...)
n.nacks[i] = nck
}
if len(n.nacks) >= maxNackCache {
copy(n.nacks, n.nacks[1:])
}
}
func (n *nackQueue) pairs(headSN uint32) ([]rtcp.NackPair, bool) {
if len(n.nacks) == 0 {
return nil, false
}
i := 0
askKF := false
var np rtcp.NackPair
var nps []rtcp.NackPair
lostIdx := -1
for _, nck := range n.nacks {
if nck.nacked >= maxNackTimes {
if nck.sn > n.kfSN {
n.kfSN = nck.sn
askKF = true
}
continue
}
if nck.sn >= headSN-2 {
n.nacks[i] = nck
i++
continue
}
n.nacks[i] = nack{
sn: nck.sn,
nacked: nck.nacked + 1,
}
i++
// first nackpair or need a new nackpair
if lostIdx < 0 || nck.sn > n.nacks[lostIdx].sn+16 {
if lostIdx >= 0 {
nps = append(nps, np)
}
np.PacketID = uint16(nck.sn)
np.LostPackets = 0
lostIdx = i - 1
continue
}
np.LostPackets |= 1 << ((nck.sn) - n.nacks[lostIdx].sn - 1)
}
// append last nackpair
if lostIdx != -1 {
nps = append(nps, np)
}
n.nacks = n.nacks[:i]
return nps, askKF
}
+196
View File
@@ -0,0 +1,196 @@
package buffer
import (
"math/rand"
"reflect"
"testing"
"time"
"github.com/pion/rtcp"
"github.com/stretchr/testify/assert"
)
func Test_nackQueue_pairs(t *testing.T) {
type fields struct {
nacks []nack
}
tests := []struct {
name string
fields fields
args []uint32
want []rtcp.NackPair
}{
{
name: "Must return correct single pairs pair",
fields: fields{
nacks: nil,
},
args: []uint32{1, 2, 4, 5},
want: []rtcp.NackPair{{
PacketID: 1,
LostPackets: 13,
}},
},
{
name: "Must return correct pair wrap",
fields: fields{
nacks: nil,
},
args: []uint32{65536, 65538, 65540, 65541, 65566, 65568}, // wrap around 65533,2,4,5
want: []rtcp.NackPair{{
PacketID: 0, // 65536
LostPackets: 1<<4 + 1<<3 + 1<<1,
},
{
PacketID: 30, // 65566
LostPackets: 1 << 1,
}},
},
{
name: "Must return 2 pairs pair",
fields: fields{
nacks: nil,
},
args: []uint32{1, 2, 4, 5, 20, 22, 24, 27},
want: []rtcp.NackPair{
{
PacketID: 1,
LostPackets: 13,
},
{
PacketID: 20,
LostPackets: 74,
},
},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
n := &nackQueue{
nacks: tt.fields.nacks,
}
for _, sn := range tt.args {
n.push(sn)
}
got, _ := n.pairs(75530)
if !reflect.DeepEqual(got, tt.want) {
t.Errorf("pairs() = %v, want %v", got, tt.want)
}
})
}
}
func Test_nackQueue_push(t *testing.T) {
type fields struct {
nacks []nack
}
type args struct {
sn []uint32
}
tests := []struct {
name string
fields fields
args args
want []uint32
}{
{
name: "Must keep packet order",
fields: fields{
nacks: make([]nack, 0, 10),
},
args: args{
sn: []uint32{3, 4, 1, 5, 8, 7, 5},
},
want: []uint32{1, 3, 4, 5, 7, 8},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
n := &nackQueue{
nacks: tt.fields.nacks,
}
for _, sn := range tt.args.sn {
n.push(sn)
}
var newSN []uint32
for _, sn := range n.nacks {
newSN = append(newSN, sn.sn)
}
assert.Equal(t, tt.want, newSN)
})
}
}
func Test_nackQueue(t *testing.T) {
type fields struct {
nacks []nack
}
type args struct {
sn []uint32
}
tests := []struct {
name string
fields fields
args args
}{
{
name: "Must keep packet order",
fields: fields{
nacks: make([]nack, 0, 10),
},
args: args{
sn: []uint32{3, 4, 1, 5, 8, 7, 5},
},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
n := nackQueue{}
r := rand.New(rand.NewSource(time.Now().UnixNano()))
for i := 0; i < 100; i++ {
assert.NotPanics(t, func() {
n.push(uint32(r.Intn(60000)))
n.remove(uint32(r.Intn(60000)))
n.pairs(60001)
})
}
})
}
}
func Test_nackQueue_remove(t *testing.T) {
type args struct {
sn []uint32
}
tests := []struct {
name string
args args
want []uint32
}{
{
name: "Must keep packet order",
args: args{
sn: []uint32{3, 4, 1, 5, 8, 7, 5},
},
want: []uint32{1, 3, 4, 7, 8},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
n := nackQueue{}
for _, sn := range tt.args.sn {
n.push(sn)
}
n.remove(5)
var newSN []uint32
for _, sn := range n.nacks {
newSN = append(newSN, sn.sn)
}
assert.Equal(t, tt.want, newSN)
})
}
}
+44
View File
@@ -0,0 +1,44 @@
package buffer
import (
"io"
"sync/atomic"
)
type RTCPReader struct {
ssrc uint32
closed atomicBool
onPacket atomic.Value //func([]byte)
onClose func()
}
func NewRTCPReader(ssrc uint32) *RTCPReader {
return &RTCPReader{ssrc: ssrc}
}
func (r *RTCPReader) Write(p []byte) (n int, err error) {
if r.closed.get() {
err = io.EOF
return
}
if f, ok := r.onPacket.Load().(func([]byte)); ok {
f(p)
}
return
}
func (r *RTCPReader) OnClose(fn func()) {
r.onClose = fn
}
func (r *RTCPReader) Close() error {
r.closed.set(true)
r.onClose()
return nil
}
func (r *RTCPReader) OnPacket(f func([]byte)) {
r.onPacket.Store(f)
}
func (r *RTCPReader) Read(_ []byte) (n int, err error) { return }