Refactor video layer selector (#1588)

* WIP commit

* WIP commit

* fix test

* FPS for VP9

* WIP commit

* test changes

* WIP commit

* h264

* codec munger

* forwarder state

* clean up a bit

* dd interface

* WIP commit

* WIP commit

* WIP commit

* WIP commit

* more TODO notes

* overshoot interface

* clean up

* clean up isTemporalSupported

* wait for key frame to resume

* clean up VP8 payload descriptor stuff

* temporal layer selector

* comment out vp9 and av1

* space

* fix test compile

* append bytes

* fix tests

* fix test
This commit is contained in:
Raja Subramanian
2023-04-08 10:57:57 +05:30
committed by GitHub
parent 57b931e9bd
commit e32eaa451f
26 changed files with 2183 additions and 1547 deletions
+30 -7
View File
@@ -10,6 +10,7 @@ import (
"github.com/gammazero/deque"
"github.com/pion/rtcp"
"github.com/pion/rtp"
"github.com/pion/rtp/codecs"
"github.com/pion/sdp/v3"
"github.com/pion/webrtc/v3"
"go.uber.org/atomic"
@@ -193,8 +194,17 @@ func (b *Buffer) Bind(params webrtc.RTPParameters, codec webrtc.RTPCodecCapabili
case strings.HasPrefix(b.mime, "video/"):
b.codecType = webrtc.RTPCodecTypeVideo
b.bucket = bucket.NewBucket(b.videoPool.Get().(*[]byte))
if b.frameRateCalculator[0] == nil && strings.EqualFold(codec.MimeType, webrtc.MimeTypeVP8) {
b.frameRateCalculator[0] = NewFrameRateCalculatorVP8(b.clockRate, b.logger)
if b.frameRateCalculator[0] == nil {
if strings.EqualFold(codec.MimeType, webrtc.MimeTypeVP8) {
b.frameRateCalculator[0] = NewFrameRateCalculatorVP8(b.clockRate, b.logger)
}
if strings.EqualFold(codec.MimeType, webrtc.MimeTypeVP9) {
frc := NewFrameRateCalculatorVP9(b.clockRate, b.logger)
for i := range b.frameRateCalculator {
b.frameRateCalculator[i] = frc.GetFrameRateCalculatorForSpatial(int32(i))
}
}
}
default:
@@ -560,12 +570,25 @@ func (b *Buffer) getExtPacket(rtpPacket *rtp.Packet, arrivalTime int64) *ExtPack
ep.Spatial = InvalidLayerSpatial // vp8 don't have spatial scalability, reset to -1
}
ep.Payload = vp8Packet
case "video/h264":
ep.KeyFrame = IsH264Keyframe(rtpPacket.Payload)
case "video/av1":
ep.KeyFrame = IsAV1Keyframe(rtpPacket.Payload)
case "video/vp9":
ep.KeyFrame = IsVP9Keyframe(rtpPacket.Payload)
if ep.DependencyDescriptor == nil {
var vp9Packet codecs.VP9Packet
_, err := vp9Packet.Unmarshal(rtpPacket.Payload)
if err != nil {
b.logger.Warnw("could not unmarshal VP9 packet", err)
return nil
}
ep.VideoLayer = VideoLayer{
Spatial: int32(vp9Packet.SID),
Temporal: int32(vp9Packet.TID),
}
ep.Payload = vp9Packet
}
ep.KeyFrame = IsVP9KeyFrame(rtpPacket.Payload)
case "video/h264":
ep.KeyFrame = IsH264KeyFrame(rtpPacket.Payload)
case "video/av1":
ep.KeyFrame = IsAV1KeyFrame(rtpPacket.Payload)
}
if ep.KeyFrame {
+143 -19
View File
@@ -4,6 +4,7 @@ import (
"container/list"
"github.com/livekit/protocol/logger"
"github.com/pion/rtp/codecs"
)
var minFramesForCalculation = [DefaultMaxLayerTemporal + 1]int{8, 15, 40}
@@ -24,8 +25,9 @@ type FrameRateCalculator interface {
}
// -----------------------------
// FrameRateCalculator based on PictureID in VP8
type FrameRateCalculatorVP8 struct {
// FrameRateCalculator based on PictureID in VPx
type frameRateCalculatorVPx struct {
frameRates [DefaultMaxLayerTemporal + 1]float32
clockRate uint32
logger logger.Logger
@@ -36,27 +38,21 @@ type FrameRateCalculatorVP8 struct {
completed bool
}
func NewFrameRateCalculatorVP8(clockRate uint32, logger logger.Logger) *FrameRateCalculatorVP8 {
return &FrameRateCalculatorVP8{
func newFrameRateCalculatorVPx(clockRate uint32, logger logger.Logger) *frameRateCalculatorVPx {
return &frameRateCalculatorVPx{
clockRate: clockRate,
logger: logger,
}
}
func (f *FrameRateCalculatorVP8) Completed() bool {
func (f *frameRateCalculatorVPx) Completed() bool {
return f.completed
}
func (f *FrameRateCalculatorVP8) RecvPacket(ep *ExtPacket) bool {
func (f *frameRateCalculatorVPx) RecvPacket(ep *ExtPacket, fn uint16) bool {
if f.completed {
return true
}
vp8, ok := ep.Payload.(VP8)
if !ok {
f.logger.Debugw("no vp8 payload", "sn", ep.Packet.SequenceNumber)
return false
}
fn := vp8.PictureID
if ep.Temporal >= int32(len(f.frameRates)) {
f.logger.Warnw("invalid temporal layer", nil, "temporal", ep.Temporal)
@@ -113,7 +109,7 @@ func (f *FrameRateCalculatorVP8) RecvPacket(ep *ExtPacket) bool {
return f.calc()
}
func (f *FrameRateCalculatorVP8) calc() bool {
func (f *frameRateCalculatorVPx) calc() bool {
var rateCounter int
for currentTemporal := int32(0); currentTemporal <= DefaultMaxLayerTemporal; currentTemporal++ {
if f.frameRates[currentTemporal] > 0 {
@@ -156,14 +152,13 @@ func (f *FrameRateCalculatorVP8) calc() bool {
if f.frameRates[2] > 0 && f.frameRates[2] > f.frameRates[1]*3 {
f.frameRates[1] = f.frameRates[2] / 2
}
f.logger.Debugw("frame rate calculated", "rate", f.frameRates)
f.reset()
return true
}
return false
}
func (f *FrameRateCalculatorVP8) reset() {
func (f *frameRateCalculatorVPx) reset() {
for i := range f.firstFrames {
f.firstFrames[i] = nil
f.secondFrames[i] = nil
@@ -175,20 +170,145 @@ func (f *FrameRateCalculatorVP8) reset() {
f.baseFrame = nil
}
func (f *FrameRateCalculatorVP8) GetFrameRate() []float32 {
func (f *frameRateCalculatorVPx) GetFrameRate() []float32 {
return f.frameRates[:]
}
// -----------------------------
// FrameRateCalculator based on Dependency descriptor
// FrameRateCalculator based on PictureID in VP8
type FrameRateCalculatorVP8 struct {
*frameRateCalculatorVPx
logger logger.Logger
}
func NewFrameRateCalculatorVP8(clockRate uint32, logger logger.Logger) *FrameRateCalculatorVP8 {
return &FrameRateCalculatorVP8{
frameRateCalculatorVPx: newFrameRateCalculatorVPx(clockRate, logger),
logger: logger,
}
}
func (f *FrameRateCalculatorVP8) RecvPacket(ep *ExtPacket) bool {
if f.frameRateCalculatorVPx.Completed() {
return true
}
vp8, ok := ep.Payload.(VP8)
if !ok {
f.logger.Debugw("no vp8 payload", "sn", ep.Packet.SequenceNumber)
return false
}
success := f.frameRateCalculatorVPx.RecvPacket(ep, vp8.PictureID)
if f.frameRateCalculatorVPx.Completed() {
f.logger.Debugw("frame rate calculated", "rate", f.frameRateCalculatorVPx.GetFrameRate())
}
return success
}
// -----------------------------
// FrameRateCalculator based on PictureID in VP9
type FrameRateCalculatorVP9 struct {
logger logger.Logger
completed bool
// VP9-TODO - this is assuming three spatial layers. As `completed` marker relies on all layers being finished, have to assume this. FIX.
// Maybe look at number of layers in livekit.TrackInfo and declare completed once advertised layers are measured
frameRateCalculatorsVPx [DefaultMaxLayerSpatial + 1]*frameRateCalculatorVPx
}
func NewFrameRateCalculatorVP9(clockRate uint32, logger logger.Logger) *FrameRateCalculatorVP9 {
f := &FrameRateCalculatorVP9{
logger: logger,
}
for i := range f.frameRateCalculatorsVPx {
f.frameRateCalculatorsVPx[i] = newFrameRateCalculatorVPx(clockRate, logger)
}
return f
}
func (f *FrameRateCalculatorVP9) Completed() bool {
return f.completed
}
func (f *FrameRateCalculatorVP9) RecvPacket(ep *ExtPacket) bool {
if f.completed {
return true
}
vp9, ok := ep.Payload.(codecs.VP9Packet)
if !ok {
f.logger.Debugw("no vp9 payload", "sn", ep.Packet.SequenceNumber)
return false
}
if ep.Spatial < 0 || ep.Spatial >= int32(len(f.frameRateCalculatorsVPx)) || f.frameRateCalculatorsVPx[ep.Spatial] == nil {
f.logger.Debugw("invalid spatial layer", "sn", ep.Packet.SequenceNumber, "spatial", ep.Spatial)
return false
}
success := f.frameRateCalculatorsVPx[ep.Spatial].RecvPacket(ep, vp9.PictureID)
completed := true
for _, frc := range f.frameRateCalculatorsVPx {
if !frc.Completed() {
completed = false
break
}
}
if completed {
f.completed = true
var frameRates [DefaultMaxLayerSpatial + 1][]float32
for i := range f.frameRateCalculatorsVPx {
frameRates[i] = f.frameRateCalculatorsVPx[i].GetFrameRate()
}
f.logger.Debugw("frame rate calculated", "rate", frameRates)
}
return success
}
func (f *FrameRateCalculatorVP9) GetFrameRateForSpatial(spatial int32) []float32 {
if spatial < 0 || spatial >= int32(len(f.frameRateCalculatorsVPx)) || f.frameRateCalculatorsVPx[spatial] == nil {
return nil
}
return f.frameRateCalculatorsVPx[spatial].GetFrameRate()
}
func (f *FrameRateCalculatorVP9) GetFrameRateCalculatorForSpatial(spatial int32) *FrameRateCalculatorForVP9Layer {
return &FrameRateCalculatorForVP9Layer{
FrameRateCalculatorVP9: f,
spatial: spatial,
}
}
// -----------------------------
type FrameRateCalculatorForVP9Layer struct {
*FrameRateCalculatorVP9
spatial int32
}
func (f *FrameRateCalculatorForVP9Layer) GetFrameRate() []float32 {
return f.FrameRateCalculatorVP9.GetFrameRateForSpatial(f.spatial)
}
// -----------------------------------------------
// FrameRateCalculator based on Dependency descriptor
type FrameRateCalculatorDD struct {
frameRates [DefaultMaxLayerSpatial + 1][DefaultMaxLayerTemporal + 1]float32
clockRate uint32
logger logger.Logger
firstFrames [DefaultMaxLayerSpatial + 1][DefaultMaxLayerTemporal + 1]*frameInfo
secondFrames [DefaultMaxLayerSpatial + 1][DefaultMaxLayerTemporal + 1]*frameInfo
spatial int
fnReceived [256]*frameInfo
baseFrame *frameInfo
completed bool
@@ -385,7 +505,7 @@ func (f *FrameRateCalculatorDD) calc() bool {
f.completed = true
f.close()
f.logger.Debugw("frame rate calculated", "spatial", f.spatial, "rate", f.frameRates)
f.logger.Debugw("frame rate calculated", "rate", f.frameRates)
return true
}
return false
@@ -424,6 +544,8 @@ func (f *FrameRateCalculatorDD) GetFrameRateCalculatorForSpatial(spatial int32)
}
}
// -----------------------------------------------
type FrameRateCalculatorForDDLayer struct {
*FrameRateCalculatorDD
spatial int32
@@ -432,3 +554,5 @@ type FrameRateCalculatorForDDLayer struct {
func (f *FrameRateCalculatorForDDLayer) GetFrameRate() []float32 {
return f.FrameRateCalculatorDD.GetFrameRateForSpatial(f.spatial)
}
// -----------------------------------------------
+128 -81
View File
@@ -4,8 +4,6 @@ import (
"encoding/binary"
"errors"
"github.com/pion/rtp/codecs"
"github.com/livekit/protocol/logger"
)
@@ -35,22 +33,23 @@ var (
*/
type VP8 struct {
FirstByte byte
S bool
PictureIDPresent int
PictureID uint16 /* 8 or 16 bits, picture ID */
MBit bool
I bool
M bool
PictureID uint16 /* 8 or 16 bits, picture ID */
TL0PICIDXPresent int
TL0PICIDX uint8 /* 8 bits temporal level zero index */
L bool
TL0PICIDX uint8 /* 8 bits temporal level zero index */
// 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
T bool
TID uint8 /* 2 bits temporal layer idx */
Y bool
KEYIDXPresent int
KEYIDX uint8 /* 5 bits of key frame idx */
K bool
KEYIDX uint8 /* 5 bits of key frame idx */
HeaderSize int
@@ -65,96 +64,94 @@ func (v *VP8) Unmarshal(payload []byte) error {
}
payloadLen := len(payload)
if payloadLen < 1 {
return errShortPacket
}
idx := 0
v.FirstByte = payload[idx]
S := payload[idx]&0x10 > 0
v.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 {
v.I = payload[idx]&0x80 > 0
v.L = payload[idx]&0x40 > 0
v.T = payload[idx]&0x20 > 0
v.K = payload[idx]&0x10 > 0
if v.L && !v.T {
return errInvalidPacket
}
// Check for PictureID
if I {
if v.I {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
v.PictureIDPresent = 1
pid := payload[idx] & 0x7f
// Check if m is 1, then Picture ID is 15 bits
if payload[idx]&0x80 > 0 {
// if m is 1, then Picture ID is 15 bits
v.M = payload[idx]&0x80 > 0
if v.M {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
v.MBit = true
v.PictureID = binary.BigEndian.Uint16([]byte{pid, payload[idx]})
} else {
v.PictureID = uint16(pid)
}
}
// Check if TL0PICIDX is present
if L {
if v.L {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
v.TL0PICIDXPresent = 1
if idx >= payloadLen {
return errShortPacket
}
v.TL0PICIDX = payload[idx]
}
if T || K {
if v.T || v.K {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
if T {
v.TIDPresent = 1
if v.T {
v.TID = (payload[idx] & 0xc0) >> 6
v.Y = (payload[idx] & 0x20) >> 5
v.Y = (payload[idx] & 0x20) > 0
}
if K {
v.KEYIDXPresent = 1
if v.K {
v.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
v.IsKeyFrame = payload[idx]&0x01 == 0 && S
v.IsKeyFrame = payload[idx]&0x01 == 0 && v.S
} else {
idx++
if payloadLen < idx+1 {
return errShortPacket
}
// Check is packet is a keyframe by looking at P bit in vp8 payload
v.IsKeyFrame = payload[idx]&0x01 == 0 && S
v.IsKeyFrame = payload[idx]&0x01 == 0 && v.S
}
v.HeaderSize = idx
return nil
}
func (v *VP8) Marshal() ([]byte, error) {
buf := make([]byte, v.HeaderSize)
err := v.MarshalTo(buf)
return buf, err
}
func (v *VP8) MarshalTo(buf []byte) error {
if len(buf) < v.HeaderSize {
return errShortPacket
@@ -162,13 +159,17 @@ func (v *VP8) MarshalTo(buf []byte) error {
idx := 0
buf[idx] = v.FirstByte
if (v.PictureIDPresent + v.TL0PICIDXPresent + v.TIDPresent + v.KEYIDXPresent) != 0 {
if v.I || v.L || v.T || v.K {
buf[idx] |= 0x80 // X bit
idx++
buf[idx] = byte(v.PictureIDPresent<<7) | byte(v.TL0PICIDXPresent<<6) | byte(v.TIDPresent<<5) | byte(v.KEYIDXPresent<<4)
xpos := idx
xval := byte(0)
idx++
if v.PictureIDPresent == 1 {
if v.MBit {
if v.I {
xval |= (1 << 7)
if v.M {
buf[idx] = 0x80 | byte((v.PictureID>>8)&0x7f)
buf[idx+1] = byte(v.PictureID & 0xff)
idx += 2
@@ -177,20 +178,31 @@ func (v *VP8) MarshalTo(buf []byte) error {
idx++
}
}
if v.TL0PICIDXPresent == 1 {
if v.L {
xval |= (1 << 6)
buf[idx] = v.TL0PICIDX
idx++
}
if v.TIDPresent == 1 || v.KEYIDXPresent == 1 {
if v.T || v.K {
buf[idx] = 0
if v.TIDPresent == 1 {
buf[idx] = v.TID<<6 | v.Y<<5
if v.T {
xval |= (1 << 5)
buf[idx] = v.TID << 6
if v.Y {
buf[idx] |= (1 << 5)
}
}
if v.KEYIDXPresent == 1 {
if v.K {
xval |= (1 << 4)
buf[idx] |= v.KEYIDX & 0x1f
}
idx++
}
buf[xpos] = xval
} else {
buf[idx] &^= 0x80 // X bit
idx++
@@ -199,7 +211,9 @@ func (v *VP8) MarshalTo(buf []byte) error {
return nil
}
func VP8PictureIdSizeDiff(mBit1 bool, mBit2 bool) int {
// -------------------------------------
func VPxPictureIdSizeDiff(mBit1 bool, mBit2 bool) int {
if mBit1 == mBit2 {
return 0
}
@@ -211,10 +225,12 @@ func VP8PictureIdSizeDiff(mBit1 bool, mBit2 bool) int {
return -1
}
// IsH264Keyframe detects if h264 payload is a keyframe
// -------------------------------------
// 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 {
func IsH264KeyFrame(payload []byte) bool {
if len(payload) < 1 {
return false
}
@@ -278,10 +294,65 @@ func IsH264Keyframe(payload []byte) bool {
return false
}
// IsAV1Keyframe detects if av1 payload is a keyframe
// -------------------------------------
func IsVP9KeyFrame(payload []byte) bool {
payloadLen := len(payload)
if payloadLen < 1 {
return false
}
idx := 0
I := payload[idx]&0x80 > 0
P := payload[idx]&0x40 > 0
L := payload[idx]&0x20 > 0
F := payload[idx]&0x10 > 0
B := payload[idx]&0x08 > 0
if F && !I {
return false
}
// Check for PictureID
if I {
idx++
if payloadLen < idx+1 {
return false
}
// Check if m is 1, then Picture ID is 15 bits
if payload[idx]&0x80 > 0 {
idx++
if payloadLen < idx+1 {
return false
}
}
}
// Check if TL0PICIDX is present
sid := -1
if L {
idx++
if payloadLen < idx+1 {
return false
}
tid := (payload[idx] >> 5) & 0x7
if !P && tid != 0 {
return false
}
sid = int((payload[idx] >> 1) & 0x7)
}
return !P && (!L || (L && sid == 0)) && B
}
// -------------------------------------
// IsAV1KeyFrame detects if av1 payload is a keyframe
// taken from https://github.com/jech/galene/blob/master/codecs/codecs.go
// all credits belongs to Juliusz Chroboczek @jech and the awesome Galene SFU
func IsAV1Keyframe(payload []byte) bool {
func IsAV1KeyFrame(payload []byte) bool {
if len(payload) < 2 {
return false
}
@@ -353,28 +424,4 @@ func IsAV1Keyframe(payload []byte) bool {
}
}
// IsVP9Keyframe detects if vp9 payload is a keyframe
// taken from https://github.com/jech/galene/blob/master/codecs/codecs.go
// all credits belongs to Juliusz Chroboczek @jech and the awesome Galene SFU
func IsVP9Keyframe(payload []byte) bool {
var vp9 codecs.VP9Packet
_, err := vp9.Unmarshal(payload)
if err != nil || len(vp9.Payload) < 1 {
return false
}
if !vp9.B {
return false
}
if (vp9.Payload[0] & 0xc0) != 0x80 {
return false
}
profile := (vp9.Payload[0] >> 4) & 0x3
if profile != 3 {
return (vp9.Payload[0] & 0xC) == 0
}
return (vp9.Payload[0] & 0x6) == 0
}
// -------------------------------------
+1 -1
View File
@@ -75,7 +75,7 @@ func TestVP8Helper_Unmarshal(t *testing.T) {
t.Errorf("Unmarshal() error = %v, wantErr %v", err, tt.wantErr)
}
if tt.checkTemporal {
require.Equal(t, tt.temporalSupport, p.TIDPresent == 1)
require.Equal(t, tt.temporalSupport, p.T)
}
if tt.checkKeyFrame {
require.Equal(t, tt.keyFrame, p.IsKeyFrame)