mirror of
https://github.com/livekit/livekit.git
synced 2026-08-21 23:09:44 +00:00
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:
co-authored by
cnderrauber
parent
289ebd32ff
commit
1e1aaeb86b
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
)
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
Reference in New Issue
Block a user