diff --git a/internal/rtpbuffer/rtpbuffer.go b/internal/rtpbuffer/rtpbuffer.go index 38121d64..93ba6daf 100644 --- a/internal/rtpbuffer/rtpbuffer.go +++ b/internal/rtpbuffer/rtpbuffer.go @@ -116,3 +116,5 @@ func (r *RTPBuffer) Get(seq uint16) *RetainablePacket { return pkt } + +func (r *RTPBuffer) Started() bool { return r.started } diff --git a/pkg/jitterbuffer/jitter_buffer.go b/pkg/jitterbuffer/jitter_buffer.go index 0741ea9a..72e08b65 100644 --- a/pkg/jitterbuffer/jitter_buffer.go +++ b/pkg/jitterbuffer/jitter_buffer.go @@ -9,6 +9,7 @@ import ( "errors" "sync" + "github.com/pion/interceptor/internal/rtpbuffer" "github.com/pion/rtp" ) @@ -66,16 +67,19 @@ type ( // order, and allows removing in either sequence number order or via a // provided timestamp. type JitterBuffer struct { - packets *PriorityQueue - minStartCount uint16 - overflowLen uint16 - lastSequence uint16 - playoutHead uint16 - playoutReady bool - state State - stats Stats - listeners map[Event][]EventListener - mutex sync.Mutex + packetFactory rtpbuffer.PacketFactory + reorderBuffer *rtpbuffer.RTPBuffer + playbackBuffer *RingBuffer + minStartCount uint16 + overflowLen uint16 + lastSequence uint16 + expectedSequence uint16 + playoutHead uint16 + playoutReady bool + state State + stats Stats + listeners map[Event][]EventListener + mutex sync.Mutex } // Stats Track interesting statistics for the life of this JitterBuffer @@ -91,14 +95,20 @@ type Stats struct { overflowCount uint32 } +var ( + // ErrInvalidOperation may be returned if a Pop or Find operation is performed on an playback buffer. + ErrInvalidOperation = errors.New("attempt to find or pop on an empty list") + // ErrNotFound will be returned if the packet cannot be found in the playblack buffer. + ErrNotFound = errors.New("packet not found") +) + // New will initialize a jitter buffer and its associated statistics. func New(opts ...Option) *JitterBuffer { jb := &JitterBuffer{ state: Buffering, stats: Stats{0, 0, 0}, minStartCount: 50, - overflowLen: 100, - packets: NewQueue(), + overflowLen: 1024, listeners: make(map[Event][]EventListener), } @@ -106,6 +116,12 @@ func New(opts ...Option) *JitterBuffer { o(jb) } + if jb.packetFactory == nil { + jb.packetFactory = rtpbuffer.NewPacketFactoryCopy() + } + jb.reorderBuffer, _ = rtpbuffer.NewRTPBuffer(jb.overflowLen) + jb.playbackBuffer = NewRingBuffer(jb.overflowLen) + return jb } @@ -117,6 +133,14 @@ func WithMinimumPacketCount(count uint16) Option { } } +// DisableCopy bypasses copy of underlying packets. It should be used when +// you are not re-using underlying buffers of packets that have been written. +func DisableCopy() Option { + return func(jb *JitterBuffer) { + jb.packetFactory = &rtpbuffer.PacketFactoryNoOp{} + } +} + // Listen will register an event listener // The jitter buffer may emit events correspnding, interested listerns should // look at Event for available events. @@ -139,12 +163,14 @@ func (jb *JitterBuffer) SetPlayoutHead(playoutHead uint16) { defer jb.mutex.Unlock() jb.playoutHead = playoutHead + jb.expectedSequence = playoutHead + jb.drain() } func (jb *JitterBuffer) updateStats(lastPktSeqNo uint16) { // If we have at least one packet, and the next packet being pushed in is not // at the expected sequence number increment the out of order count - if jb.packets.Length() > 0 && lastPktSeqNo != (jb.lastSequence+1) { + if jb.reorderBuffer.Started() && lastPktSeqNo != (jb.lastSequence+1) { jb.stats.outOfOrderCount++ } jb.lastSequence = lastPktSeqNo @@ -157,21 +183,24 @@ func (jb *JitterBuffer) Push(packet *rtp.Packet) { jb.mutex.Lock() defer jb.mutex.Unlock() - if jb.packets.Length() == 0 { - jb.emit(StartBuffering) + rPacket, err := jb.packetFactory.NewPacket(&packet.Header, packet.Payload, 0, 0) + if err != nil { + return } - if jb.packets.Length() > jb.overflowLen { - jb.stats.overflowCount++ - jb.emit(BufferOverflow) + if jb.playbackBuffer.Length() == 0 { + jb.emit(StartBuffering) } - if !jb.playoutReady && jb.packets.Length() == 0 { + if !jb.reorderBuffer.Started() { jb.playoutHead = packet.SequenceNumber + jb.expectedSequence = packet.SequenceNumber } - jb.updateStats(packet.SequenceNumber) - jb.packets.Push(packet, packet.SequenceNumber) + + jb.reorderBuffer.Add(rPacket) + jb.drain() + jb.updateState() } @@ -183,7 +212,7 @@ func (jb *JitterBuffer) emit(event Event) { func (jb *JitterBuffer) updateState() { // For now, we only look at the number of packets captured in the play buffer - if jb.packets.Length() >= jb.minStartCount && jb.state == Buffering { + if jb.playbackBuffer.Length() >= jb.minStartCount && jb.state == Buffering { jb.state = Emitting jb.playoutReady = true jb.emit(BeginPlayback) @@ -200,29 +229,52 @@ func (jb *JitterBuffer) updateState() { func (jb *JitterBuffer) Peek(playoutHead bool) (*rtp.Packet, error) { jb.mutex.Lock() defer jb.mutex.Unlock() - if jb.packets.Length() < 1 { + if !jb.reorderBuffer.Started() { return nil, ErrBufferUnderrun } + + var packet *rtpbuffer.RetainablePacket if playoutHead && jb.state == Emitting { - return jb.packets.Find(jb.playoutHead) + packet = jb.playbackBuffer.Peek() + if packet != nil { + if err := packet.Retain(); err != nil { + return nil, ErrNotFound + } + } + } else { + packet = jb.reorderBuffer.Get(jb.lastSequence) } - return jb.packets.Find(jb.lastSequence) + if packet == nil { + return nil, ErrNotFound + } + + return jb.takePacket(packet), nil } // Pop an RTP packet from the jitter buffer at the current playout head. func (jb *JitterBuffer) Pop() (*rtp.Packet, error) { + packet, err := jb.popRetainable() + if err != nil { + return nil, err + } + + return jb.takePacket(packet), nil +} + +// Same as Pop, except it returns a RetainablePacket. +func (jb *JitterBuffer) popRetainable() (*rtpbuffer.RetainablePacket, error) { jb.mutex.Lock() defer jb.mutex.Unlock() if jb.state != Emitting { return nil, ErrPopWhileBuffering } - packet, err := jb.packets.PopAt(jb.playoutHead) - if err != nil { + packet := jb.playbackBuffer.Pop() + if packet == nil { jb.stats.underflowCount++ jb.emit(BufferUnderflow) - return nil, err + return nil, ErrNotFound } jb.playoutHead = (jb.playoutHead + 1) jb.updateState() @@ -237,17 +289,17 @@ func (jb *JitterBuffer) PopAtSequence(sq uint16) (*rtp.Packet, error) { if jb.state != Emitting { return nil, ErrPopWhileBuffering } - packet, err := jb.packets.PopAt(sq) - if err != nil { + packet := jb.playbackBuffer.PopAt(sq) + if packet == nil { jb.stats.underflowCount++ jb.emit(BufferUnderflow) - return nil, err + return nil, ErrNotFound } - jb.playoutHead = (jb.playoutHead + 1) + jb.playoutHead = sq + 1 jb.updateState() - return packet, nil + return jb.takePacket(packet), nil } // PeekAtSequence will return an RTP packet from the jitter buffer at the specified Sequence @@ -255,12 +307,12 @@ func (jb *JitterBuffer) PopAtSequence(sq uint16) (*rtp.Packet, error) { func (jb *JitterBuffer) PeekAtSequence(sq uint16) (*rtp.Packet, error) { jb.mutex.Lock() defer jb.mutex.Unlock() - packet, err := jb.packets.Find(sq) - if err != nil { - return nil, err + packet := jb.reorderBuffer.Get(sq) + if packet == nil { + return nil, ErrNotFound } - return packet, nil + return jb.takePacket(packet), nil } // PopAtTimestamp pops an RTP packet from the jitter buffer with the provided timestamp @@ -271,26 +323,65 @@ func (jb *JitterBuffer) PopAtTimestamp(ts uint32) (*rtp.Packet, error) { if jb.state != Emitting { return nil, ErrPopWhileBuffering } - packet, err := jb.packets.PopAtTimestamp(ts) - if err != nil { + packet := jb.playbackBuffer.PopAtTimestamp(ts) + if packet == nil { jb.stats.underflowCount++ jb.emit(BufferUnderflow) - return nil, err + return nil, ErrNotFound } + jb.playoutHead = packet.Header().SequenceNumber + 1 jb.updateState() - return packet, nil + return jb.takePacket(packet), nil +} + +// Unwrap the packet. +func (jb *JitterBuffer) takePacket(rPacket *rtpbuffer.RetainablePacket) *rtp.Packet { + header := *rPacket.Header() + payload := rPacket.Payload() + if _, ok := jb.packetFactory.(*rtpbuffer.PacketFactoryNoOp); !ok { + out := make([]byte, len(payload)) + copy(out, payload) + payload = out + } + rPacket.Release() + + return &rtp.Packet{ + Header: header, + Payload: payload, + } +} + +// Move the packets into the playback buffer as long as they are in order. +func (jb *JitterBuffer) drain() { + for range jb.overflowLen { + rPacket := jb.reorderBuffer.Get(jb.expectedSequence) + if rPacket == nil { + break + } + if !jb.playbackBuffer.Push(rPacket) { + rPacket.Release() + jb.stats.overflowCount++ + jb.emit(BufferOverflow) + break + } + jb.expectedSequence++ + } } // Clear will empty the buffer and optionally reset the state. func (jb *JitterBuffer) Clear(resetState bool) { jb.mutex.Lock() defer jb.mutex.Unlock() - jb.packets.Clear() + jb.reorderBuffer.Clear() + jb.playbackBuffer.Clear() if resetState { jb.lastSequence = 0 + jb.expectedSequence = 0 + jb.playoutHead = 0 + jb.playoutReady = false jb.state = Buffering jb.stats = Stats{0, 0, 0} jb.minStartCount = 50 diff --git a/pkg/jitterbuffer/jitter_buffer_test.go b/pkg/jitterbuffer/jitter_buffer_test.go index 8894cf8f..2c94895a 100644 --- a/pkg/jitterbuffer/jitter_buffer_test.go +++ b/pkg/jitterbuffer/jitter_buffer_test.go @@ -27,7 +27,7 @@ func TestJitterBuffer(t *testing.T) { jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 5012, Timestamp: 512}, Payload: []byte{0x02}}) assert.Equal(jb.stats.outOfOrderCount, uint32(1)) - assert.Equal(jb.packets.Length(), uint16(4)) + assert.Equal(jb.playbackBuffer.Length(), uint16(3)) assert.Equal(jb.lastSequence, uint16(5012)) }) t.Run("Appends packets and wraps", func(*testing.T) { @@ -39,7 +39,7 @@ func TestJitterBuffer(t *testing.T) { jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 0, Timestamp: 512}, Payload: []byte{0x02}}) - assert.Equal(jb.packets.Length(), uint16(2)) + assert.Equal(jb.playbackBuffer.Length(), uint16(2)) assert.Equal(jb.lastSequence, uint16(0)) head, err := jb.Pop() @@ -63,7 +63,7 @@ func TestJitterBuffer(t *testing.T) { }, ) } - assert.Equal(jb.packets.Length(), uint16(100)) + assert.Equal(jb.playbackBuffer.Length(), uint16(100)) assert.Equal(jb.state, Emitting) assert.Equal(jb.playoutHead, uint16(5012)) head, err := jb.Pop() @@ -88,7 +88,7 @@ func TestJitterBuffer(t *testing.T) { }, ) } - assert.Equal(jb.packets.Length(), uint16(2)) + assert.Equal(jb.playbackBuffer.Length(), uint16(2)) assert.Equal(jb.state, Emitting) assert.Equal(jb.playoutHead, uint16(5012)) head, err := jb.Pop() @@ -105,7 +105,7 @@ func TestJitterBuffer(t *testing.T) { //nolint:gosec // G115 jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: sqnum, Timestamp: uint32(512 + i)}, Payload: []byte{0x02}}) } - assert.Equal(jb.packets.Length(), uint16(100)) + assert.Equal(jb.playbackBuffer.Length(), uint16(100)) assert.Equal(jb.state, Emitting) assert.Equal(jb.playoutHead, uint16(math.MaxUint16-32)) head, err := jb.Pop() @@ -129,35 +129,40 @@ func TestJitterBuffer(t *testing.T) { //nolint:gosec // G115 jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: sqnum, Timestamp: uint32(512 + i)}, Payload: []byte{0x02}}) } - assert.Equal(jb.packets.Length(), uint16(100)) + assert.Equal(jb.playbackBuffer.Length(), uint16(100)) assert.Equal(jb.state, Emitting) + // this will discard (one) ts=512 because it is older + // and return the oldest ts=513 head, err := jb.PopAtTimestamp(uint32(513)) assert.Equal(head.SequenceNumber, uint16(math.MaxUint16-32+1)) assert.Equal(err, nil) - head, err = jb.PopAtTimestamp(uint32(513)) - assert.Equal(head, (*rtp.Packet)(nil)) + head, err = jb.PopAtTimestamp(uint32(513)) // 32+1 + assert.Equal(head, (*rtp.Packet)(nil)) // there is only one ts=513 assert.NotEqual(err, nil) + // PopAtTimestamp establishes a new playoutHead + assert.Equal(jb.playoutHead, uint16(math.MaxUint16-32+2)) head, err = jb.Pop() - assert.Equal(head.SequenceNumber, uint16(math.MaxUint16-32)) + assert.Equal(head.SequenceNumber, uint16(math.MaxUint16-32+2)) assert.Equal(err, nil) }) t.Run("Can peek at a packet", func(*testing.T) { jb := New() + for i := range 100 { + sqnum := uint16((math.MaxUint16 - 32 + i)) //nolint:gosec // G115 + //nolint:gosec // G115 + jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: sqnum, Timestamp: uint32(512 + i)}, Payload: []byte{0x02}}) + } jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}}) jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 5001, Timestamp: 501}, Payload: []byte{0x02}}) jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 5002, Timestamp: 502}, Payload: []byte{0x02}}) pkt, err := jb.Peek(false) assert.Equal(pkt.SequenceNumber, uint16(5002)) assert.Equal(err, nil) - for i := range 100 { - sqnum := uint16((math.MaxUint16 - 32 + i)) //nolint:gosec // G115 - //nolint:gosec // G115 - jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: sqnum, Timestamp: uint32(512 + i)}, Payload: []byte{0x02}}) - } pkt, err = jb.Peek(true) - assert.Equal(pkt.SequenceNumber, uint16(5000)) + + assert.Equal(pkt.SequenceNumber, uint16(math.MaxUint16-32)) assert.Equal(err, nil) }) @@ -170,7 +175,7 @@ func TestJitterBuffer(t *testing.T) { } jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1019, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1020, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) - assert.Equal(jb.packets.Length(), uint16(52)) + assert.Equal(jb.playbackBuffer.Length(), uint16(50)) assert.Equal(jb.state, Emitting) head, err := jb.PopAtSequence(uint16(9000)) assert.Equal(head, (*rtp.Packet)(nil)) @@ -184,19 +189,20 @@ func TestJitterBuffer(t *testing.T) { //nolint:gosec // G115 jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: sqnum, Timestamp: uint32(512 + i)}, Payload: []byte{0x02}}) } - jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1019, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) - jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1020, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) - assert.Equal(jb.packets.Length(), uint16(52)) + jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 17, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) + jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 18, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) + jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 19, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) + assert.Equal(jb.playbackBuffer.Length(), uint16(53)) assert.Equal(jb.state, Emitting) head, err := jb.PopAtTimestamp(uint32(9000)) - assert.Equal(head.SequenceNumber, uint16(1019)) + assert.Equal(head.SequenceNumber, uint16(17)) assert.Equal(err, nil) head, err = jb.PopAtTimestamp(uint32(9000)) - assert.Equal(head.SequenceNumber, uint16(1020)) + assert.Equal(head.SequenceNumber, uint16(18)) assert.Equal(err, nil) head, err = jb.Pop() - assert.Equal(head.SequenceNumber, uint16(math.MaxUint16-32)) + assert.Equal(head.SequenceNumber, uint16(19)) assert.Equal(err, nil) }) @@ -209,7 +215,7 @@ func TestJitterBuffer(t *testing.T) { } jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1019, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) jb.Push(&rtp.Packet{Header: rtp.Header{SequenceNumber: 1020, Timestamp: uint32(9000)}, Payload: []byte{0x02}}) - assert.Equal(jb.packets.Length(), uint16(52)) + assert.Equal(jb.playbackBuffer.Length(), uint16(50)) assert.Equal(jb.state, Emitting) head, err := jb.PeekAtSequence(uint16(1019)) assert.Equal(head.SequenceNumber, uint16(1019)) @@ -276,6 +282,6 @@ func TestJitterBuffer(t *testing.T) { jb.Clear(true) assert.Equal(jb.lastSequence, uint16(0)) assert.Equal(jb.stats.outOfOrderCount, uint32(0)) - assert.Equal(jb.packets.Length(), uint16(0)) + assert.Equal(jb.playbackBuffer.Length(), uint16(0)) }) } diff --git a/pkg/jitterbuffer/option.go b/pkg/jitterbuffer/option.go index a6dd4a5d..0da46df6 100644 --- a/pkg/jitterbuffer/option.go +++ b/pkg/jitterbuffer/option.go @@ -27,3 +27,12 @@ func WithLoggerFactory(loggerFactory logging.LoggerFactory) ReceiverInterceptorO return nil } } + +// WithBufferOptions configures JitterBuffers created for each remote stream. +func WithBufferOptions(opts ...Option) ReceiverInterceptorOption { + return func(d *ReceiverInterceptor) error { + d.bufferOptions = append(d.bufferOptions, opts...) + + return nil + } +} diff --git a/pkg/jitterbuffer/priority_queue.go b/pkg/jitterbuffer/priority_queue.go deleted file mode 100644 index 26bdcf47..00000000 --- a/pkg/jitterbuffer/priority_queue.go +++ /dev/null @@ -1,203 +0,0 @@ -// SPDX-FileCopyrightText: 2026 The Pion community -// SPDX-License-Identifier: MIT - -package jitterbuffer - -import ( - "errors" - - "github.com/pion/rtp" -) - -// PriorityQueue provides a linked list sorting of RTP packets by SequenceNumber. -type PriorityQueue struct { - next *node - length uint16 -} - -type node struct { - val *rtp.Packet - next *node - prev *node - priority uint16 -} - -var ( - // ErrInvalidOperation may be returned if a Pop or Find operation is performed on an empty queue. - ErrInvalidOperation = errors.New("attempt to find or pop on an empty list") - // ErrNotFound will be returned if the packet cannot be found in the queue. - ErrNotFound = errors.New("priority not found") -) - -// NewQueue will create a new PriorityQueue whose order relies on monotonically -// increasing Sequence Number, wrapping at MaxUint16, so -// a packet with sequence number MaxUint16 - 1 will be after 0. -func NewQueue() *PriorityQueue { - return &PriorityQueue{ - next: nil, - length: 0, - } -} - -func newNode(val *rtp.Packet, priority uint16) *node { - return &node{ - val: val, - prev: nil, - next: nil, - priority: priority, - } -} - -// Find a packet in the queue with the provided sequence number, -// regardless of position (the packet is retained in the queue). -func (q *PriorityQueue) Find(sqNum uint16) (*rtp.Packet, error) { - next := q.next - for next != nil { - if next.priority == sqNum { - return next.val, nil - } - next = next.next - } - - return nil, ErrNotFound -} - -// Push will insert a packet in to the queue in order of sequence number. -func (q *PriorityQueue) Push(val *rtp.Packet, priority uint16) { - newPq := newNode(val, priority) - if q.next == nil { - q.next = newPq - q.length++ - - return - } - if priority < q.next.priority { - newPq.next = q.next - q.next.prev = newPq - q.next = newPq - q.length++ - - return - } - head := q.next - prev := q.next - for head != nil { - if priority <= head.priority { - break - } - prev = head - head = head.next - } - if head == nil { - if prev != nil { - prev.next = newPq - } - newPq.prev = prev - } else { - newPq.next = head - newPq.prev = prev - if prev != nil { - prev.next = newPq - } - head.prev = newPq - } - q.length++ -} - -// Length will get the total length of the queue. -func (q *PriorityQueue) Length() uint16 { - return q.length -} - -// Pop removes the first element from the queue, regardless -// sequence number. -func (q *PriorityQueue) Pop() (*rtp.Packet, error) { - if q.next == nil { - return nil, ErrInvalidOperation - } - val := q.next.val - q.next.val = nil - q.length-- - q.next = q.next.next - - return val, nil -} - -// PopAt removes an element at the specified sequence number (priority). -func (q *PriorityQueue) PopAt(sqNum uint16) (*rtp.Packet, error) { - if q.next == nil { - return nil, ErrInvalidOperation - } - if q.next.priority == sqNum { - val := q.next.val - q.next.val = nil - q.next = q.next.next - q.length-- - - return val, nil - } - pos := q.next - prev := q.next.prev - for pos != nil { - if pos.priority == sqNum { - val := pos.val - pos.val = nil - prev.next = pos.next - if prev.next != nil { - prev.next.prev = prev - } - q.length-- - - return val, nil - } - prev = pos - pos = pos.next - } - - return nil, ErrNotFound -} - -// PopAtTimestamp removes and returns a packet at the given RTP Timestamp, regardless -// sequence number order. -func (q *PriorityQueue) PopAtTimestamp(timestamp uint32) (*rtp.Packet, error) { - if q.next == nil { - return nil, ErrInvalidOperation - } - if q.next.val.Timestamp == timestamp { - val := q.next.val - q.next.val = nil - q.next = q.next.next - q.length-- - - return val, nil - } - pos := q.next - prev := q.next.prev - for pos != nil { - if pos.val.Timestamp == timestamp { - val := pos.val - pos.val = nil - prev.next = pos.next - if prev.next != nil { - prev.next.prev = prev - } - q.length-- - - return val, nil - } - prev = pos - pos = pos.next - } - - return nil, ErrNotFound -} - -// Clear will empty a PriorityQueue. -func (q *PriorityQueue) Clear() { - next := q.next - q.length = 0 - for next != nil { - next.prev = nil - next = next.next - } -} diff --git a/pkg/jitterbuffer/priority_queue_test.go b/pkg/jitterbuffer/priority_queue_test.go deleted file mode 100644 index 40d2cedb..00000000 --- a/pkg/jitterbuffer/priority_queue_test.go +++ /dev/null @@ -1,230 +0,0 @@ -// SPDX-FileCopyrightText: 2026 The Pion community -// SPDX-License-Identifier: MIT - -package jitterbuffer - -import ( - "runtime" - "sync/atomic" - "testing" - "time" - - "github.com/pion/rtp" - "github.com/stretchr/testify/assert" -) - -func TestPriorityQueue(t *testing.T) { - assert := assert.New(t) - - t.Run("Appends packets in order", func(*testing.T) { - pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} - q := NewQueue() - q.Push(pkt, pkt.SequenceNumber) - pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5004, Timestamp: 500}, Payload: []byte{0x02}} - q.Push(pkt2, pkt2.SequenceNumber) - assert.Equal(q.next.next.val, pkt2) - assert.Equal(q.next.priority, uint16(5000)) - assert.Equal(q.next.next.priority, uint16(5004)) - }) - - t.Run("Appends many in order", func(*testing.T) { - queue := NewQueue() - for i := range 100 { - //nolint:gosec // G115 - queue.Push( - &rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: uint16(5012 + i), - Timestamp: uint32(512 + i), - }, - Payload: []byte{0x02}, - }, - uint16(5012+i), - ) - } - assert.Equal(uint16(100), queue.Length()) - last := (*node)(nil) - cur := queue.next - for cur != nil { - last = cur - cur = cur.next - if cur != nil { - assert.Equal(cur.priority, last.priority+1) - } - } - assert.Equal(queue.next.priority, uint16(5012)) - assert.Equal(last.priority, uint16(5012+99)) - }) - - t.Run("Can remove an element", func(*testing.T) { - pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} - queue := NewQueue() - queue.Push(pkt, pkt.SequenceNumber) - pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5004, Timestamp: 500}, Payload: []byte{0x02}} - queue.Push(pkt2, pkt2.SequenceNumber) - for i := range 100 { - //nolint:gosec // G115 - queue.Push( - &rtp.Packet{ - Header: rtp.Header{SequenceNumber: uint16(5012 + i), Timestamp: uint32(512 + i)}, - Payload: []byte{0x02}, - }, - uint16(5012+i), - ) - } - popped, _ := queue.Pop() - assert.Equal(popped.SequenceNumber, uint16(5000)) - _, _ = queue.Pop() - nextPop, _ := queue.Pop() - assert.Equal(nextPop.SequenceNumber, uint16(5012)) - }) - - t.Run("Appends in order", func(*testing.T) { - queue := NewQueue() - for i := range 100 { - queue.Push( - &rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: uint16(5012 + i), //nolint:gosec // G115 - Timestamp: uint32(512 + i), //nolint:gosec // G115 - }, - Payload: []byte{0x02}, - }, - uint16(5012+i), //nolint:gosec // G115 - ) - } - assert.Equal(uint16(100), queue.Length()) - pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} - queue.Push(pkt, pkt.SequenceNumber) - assert.Equal(pkt, queue.next.val) - assert.Equal(uint16(101), queue.Length()) - assert.Equal(queue.next.priority, uint16(5000)) - }) - - t.Run("Can find", func(*testing.T) { - queue := NewQueue() - for i := range 100 { - //nolint:gosec // G115 - queue.Push( - &rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: uint16(5012 + i), - Timestamp: uint32(512 + i), - }, - Payload: []byte{0x02}, - }, - uint16(5012+i), - ) - } - pkt, err := queue.Find(5012) - assert.Equal(pkt.SequenceNumber, uint16(5012)) - assert.Equal(err, nil) - }) - - t.Run("Updates the length when PopAt* are called", func(*testing.T) { - pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} - queue := NewQueue() - queue.Push(pkt, pkt.SequenceNumber) - pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5004, Timestamp: 500}, Payload: []byte{0x02}} - queue.Push(pkt2, pkt2.SequenceNumber) - for i := range 100 { - //nolint:gosec // G115 - queue.Push( - &rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: uint16(5012 + i), - Timestamp: uint32(512 + i), - }, - Payload: []byte{0x02}, - }, - uint16(5012+i), - ) - } - assert.Equal(uint16(102), queue.Length()) - popped, _ := queue.PopAt(uint16(5012)) - assert.Equal(popped.SequenceNumber, uint16(5012)) - assert.Equal(uint16(101), queue.Length()) - - popped, err := queue.PopAtTimestamp(uint32(500)) - assert.Equal(popped.SequenceNumber, uint16(5000)) - assert.Equal(uint16(100), queue.Length()) - assert.Equal(err, nil) - }) -} - -func TestPriorityQueue_Find(t *testing.T) { - packets := NewQueue() - - packets.Push(&rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: 1000, - Timestamp: 5, - SSRC: 5, - }, - Payload: []uint8{0xA}, - }, 1000) - - _, err := packets.PopAt(1000) - assert.NoError(t, err) - - _, err = packets.Find(1001) - assert.Error(t, err) -} - -func TestPriorityQueue_Clean(t *testing.T) { - packets := NewQueue() - packets.Clear() - packets.Push(&rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: 1000, - Timestamp: 5, - SSRC: 5, - }, - Payload: []uint8{0xA}, - }, 1000) - assert.EqualValues(t, 1, packets.Length()) - packets.Clear() -} - -func TestPriorityQueue_Unreference(t *testing.T) { - packets := NewQueue() - - var refs int64 - finalizer := func(*rtp.Packet) { - atomic.AddInt64(&refs, -1) - } - - numPkts := 100 - for i := range numPkts { - atomic.AddInt64(&refs, 1) - seq := uint16(i) //nolint:gosec // G115 - p := rtp.Packet{ - Header: rtp.Header{ - SequenceNumber: seq, - Timestamp: uint32(i + 42), //nolint:gosec // G115 - }, - Payload: []byte{byte(i)}, - } - runtime.SetFinalizer(&p, finalizer) - packets.Push(&p, seq) - } - for i := 0; i < numPkts-1; i++ { - switch i % 3 { - case 0: - packets.Pop() //nolint - case 1: - packets.PopAt(uint16(i)) //nolint - case 2: - packets.PopAtTimestamp(uint32(i + 42)) //nolint - } - } - - runtime.GC() - time.Sleep(10 * time.Millisecond) - - remainedRefs := atomic.LoadInt64(&refs) - runtime.KeepAlive(packets) - - // only the last packet should be still referenced - assert.Equal(t, int64(1), remainedRefs) -} diff --git a/pkg/jitterbuffer/receiver_interceptor.go b/pkg/jitterbuffer/receiver_interceptor.go index 657282bf..9961b4a8 100644 --- a/pkg/jitterbuffer/receiver_interceptor.go +++ b/pkg/jitterbuffer/receiver_interceptor.go @@ -19,8 +19,8 @@ type InterceptorFactory struct { // NewInterceptor constructs a new ReceiverInterceptor. func (g *InterceptorFactory) NewInterceptor(_ string) (interceptor.Interceptor, error) { receiverInterceptor := &ReceiverInterceptor{ - close: make(chan struct{}), - buffer: New(), + close: make(chan struct{}), + buffers: make(map[uint32]*JitterBuffer), } for _, opt := range g.opts { @@ -58,7 +58,8 @@ func (g *InterceptorFactory) NewInterceptor(_ string) (interceptor.Interceptor, // arriving) quickly enough. type ReceiverInterceptor struct { interceptor.NoOp - buffer *JitterBuffer + buffers map[uint32]*JitterBuffer + bufferOptions []Option m sync.Mutex wg sync.WaitGroup close chan struct{} @@ -74,8 +75,14 @@ func NewInterceptor(opts ...ReceiverInterceptorOption) (*InterceptorFactory, err // BindRemoteStream lets you modify any incoming RTP packets. It is called once for per RemoteStream. // The returned method will be called once per rtp packet. func (i *ReceiverInterceptor) BindRemoteStream( - _ *interceptor.StreamInfo, reader interceptor.RTPReader, + info *interceptor.StreamInfo, reader interceptor.RTPReader, ) interceptor.RTPReader { + buffer := New(i.bufferOptions...) + + i.m.Lock() + i.buffers[info.SSRC] = buffer + i.m.Unlock() + return interceptor.RTPReaderFunc(func(b []byte, a interceptor.Attributes) (int, interceptor.Attributes, error) { buf := make([]byte, len(b)) n, attr, err := reader.Read(buf, a) @@ -83,32 +90,36 @@ func (i *ReceiverInterceptor) BindRemoteStream( return n, attr, err } packet := &rtp.Packet{} - if err := packet.Unmarshal(buf); err != nil { + if err := packet.Unmarshal(buf[:n]); err != nil { return 0, nil, err } i.m.Lock() defer i.m.Unlock() - i.buffer.Push(packet) - if i.buffer.state == Emitting { - newPkt, err := i.buffer.Pop() - if err != nil { - return 0, nil, err - } - nlen, err := newPkt.MarshalTo(b) - - return nlen, attr, err + buffer.Push(packet) + rPacket, err := buffer.popRetainable() + if err != nil { + return 0, attr, ErrPopWhileBuffering } + out := rtp.Packet{ + Header: *rPacket.Header(), + Payload: rPacket.Payload(), + } + nlen, err := out.MarshalTo(b) + rPacket.Release() - return n, attr, ErrPopWhileBuffering + return nlen, attr, err }) } // UnbindRemoteStream is called when the Stream is removed. It can be used to clean up any data related to that track. -func (i *ReceiverInterceptor) UnbindRemoteStream(_ *interceptor.StreamInfo) { +func (i *ReceiverInterceptor) UnbindRemoteStream(info *interceptor.StreamInfo) { defer i.wg.Wait() i.m.Lock() defer i.m.Unlock() - i.buffer.Clear(true) + if buffer, ok := i.buffers[info.SSRC]; ok { + buffer.Clear(true) + delete(i.buffers, info.SSRC) + } } // Close closes the interceptor. @@ -116,7 +127,10 @@ func (i *ReceiverInterceptor) Close() error { defer i.wg.Wait() i.m.Lock() defer i.m.Unlock() - i.buffer.Clear(true) + for ssrc, buffer := range i.buffers { + buffer.Clear(true) + delete(i.buffers, ssrc) + } return nil } diff --git a/pkg/jitterbuffer/ring_buffer.go b/pkg/jitterbuffer/ring_buffer.go new file mode 100644 index 00000000..cd68c5b1 --- /dev/null +++ b/pkg/jitterbuffer/ring_buffer.go @@ -0,0 +1,101 @@ +// SPDX-FileCopyrightText: 2026 The Pion community +// SPDX-License-Identifier: MIT + +// RingBuffer is a classic Ring Buffer. +package jitterbuffer + +import ( + "github.com/pion/interceptor/internal/rtpbuffer" +) + +type RingBuffer struct { + buffer []*rtpbuffer.RetainablePacket + read, write, length uint16 +} + +func NewRingBuffer(size uint16) *RingBuffer { + return &RingBuffer{buffer: make([]*rtpbuffer.RetainablePacket, size)} +} + +func (r *RingBuffer) Push(rPacket *rtpbuffer.RetainablePacket) bool { + if r.Full() { + return false + } + r.buffer[r.write] = rPacket + r.write = (r.write + 1) % uint16(len(r.buffer)) + r.length++ + return true +} + +func (r *RingBuffer) Pop() *rtpbuffer.RetainablePacket { + if r.Empty() { + return nil + } + + rPacket := r.buffer[r.read] + r.read = (r.read + 1) % uint16(len(r.buffer)) + r.length-- + return rPacket +} + +func (r *RingBuffer) Peek() *rtpbuffer.RetainablePacket { + if r.Empty() { + return nil + } + + rPacket := r.buffer[r.read] + return rPacket +} + +func (r *RingBuffer) PopAt(sequenceNumber uint16) *rtpbuffer.RetainablePacket { + return r.popMatching(func(rPacket *rtpbuffer.RetainablePacket) bool { + return sequenceNumber == rPacket.Header().SequenceNumber + }) +} + +func (r *RingBuffer) PopAtTimestamp(timestamp uint32) *rtpbuffer.RetainablePacket { + return r.popMatching(func(rPacket *rtpbuffer.RetainablePacket) bool { + return timestamp == rPacket.Header().Timestamp + }) +} + +// popMatching removes and returns the first packet matched, discards the packets ahead of it, and handing ownership +func (r *RingBuffer) popMatching(match func(*rtpbuffer.RetainablePacket) bool) *rtpbuffer.RetainablePacket { + var extra uint16 + found := false + for i := range r.length { + rPacket := r.buffer[(r.read+i)%uint16(len(r.buffer))] + if match(rPacket) { + extra = i + found = true + break + } + } + + if !found { + return nil + } + + for range extra { + r.Pop().Release() + } + + return r.Pop() +} + +func (r *RingBuffer) Clear() { + for i := range r.length { + idx := (r.read + i) % uint16(len(r.buffer)) + if pkt := r.buffer[idx]; pkt != nil { + pkt.Release() + r.buffer[idx] = nil + } + } + r.read = 0 + r.write = 0 + r.length = 0 +} + +func (r *RingBuffer) Full() bool { return r.length == uint16(len(r.buffer)) } +func (r *RingBuffer) Empty() bool { return r.length == 0 } +func (r *RingBuffer) Length() uint16 { return r.length } diff --git a/pkg/jitterbuffer/ring_buffer_test.go b/pkg/jitterbuffer/ring_buffer_test.go new file mode 100644 index 00000000..d2edb4c3 --- /dev/null +++ b/pkg/jitterbuffer/ring_buffer_test.go @@ -0,0 +1,269 @@ +// SPDX-FileCopyrightText: 2026 The Pion community +// SPDX-License-Identifier: MIT + +package jitterbuffer + +import ( + "runtime" + "sync/atomic" + "testing" + "time" + + "github.com/pion/interceptor/internal/rtpbuffer" + "github.com/pion/rtp" + "github.com/stretchr/testify/assert" +) + +func TestRingBuffer(t *testing.T) { + assert := assert.New(t) + + t.Run("Appends packets in order", func(t *testing.T) { + q := NewRingBuffer(16) + pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(q.Push(buildRetainablePacket(t, pkt))) + pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5001, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(q.Push(buildRetainablePacket(t, pkt2))) + assert.Equal(uint16(5000), q.Peek().Header().SequenceNumber) + assert.Equal(uint16(2), q.Length()) + assert.Equal(uint16(5000), q.Pop().Header().SequenceNumber) + assert.Equal(uint16(5001), q.Peek().Header().SequenceNumber) + }) + + t.Run("Appends many in order", func(t *testing.T) { + queue := NewRingBuffer(128) + for i := range 100 { + assert.True(queue.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: uint16(5012 + i), + Timestamp: uint32(512 + i), + }, + Payload: []byte{0x02}, + }))) + } + assert.Equal(uint16(100), queue.Length()) + prev := queue.Peek().Header().SequenceNumber + assert.Equal(uint16(5012), prev) + for range 99 { + popped := queue.Pop() + assert.NotNil(popped) + assert.Equal(prev, popped.Header().SequenceNumber) + assert.Equal(prev+1, queue.Peek().Header().SequenceNumber) + prev = queue.Peek().Header().SequenceNumber + } + assert.Equal(uint16(5012+99), queue.Peek().Header().SequenceNumber) + }) + + t.Run("Can remove an element", func(t *testing.T) { + queue := NewRingBuffer(128) + pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(queue.Push(buildRetainablePacket(t, pkt))) + pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5001, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(queue.Push(buildRetainablePacket(t, pkt2))) + for i := range 100 { + assert.True(queue.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: uint16(5002 + i), + Timestamp: uint32(512 + i), + }, + Payload: []byte{0x02}, + }))) + } + popped := queue.Pop() + assert.Equal(uint16(5000), popped.Header().SequenceNumber) + _ = queue.Pop() + nextPop := queue.Pop() + assert.Equal(uint16(5002), nextPop.Header().SequenceNumber) + }) + + t.Run("Appends at end", func(t *testing.T) { + queue := NewRingBuffer(128) + for i := range 100 { + assert.True(queue.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: uint16(5012 + i), + Timestamp: uint32(512 + i), + }, + Payload: []byte{0x02}, + }))) + } + assert.Equal(uint16(100), queue.Length()) + pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(queue.Push(buildRetainablePacket(t, pkt))) + assert.Equal(uint16(5012), queue.Peek().Header().SequenceNumber) + assert.Equal(uint16(101), queue.Length()) + }) + + t.Run("Can peek", func(t *testing.T) { + queue := NewRingBuffer(128) + for i := range 100 { + assert.True(queue.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: uint16(5012 + i), + Timestamp: uint32(512 + i), + }, + Payload: []byte{0x02}, + }))) + } + pkt := queue.Peek() + assert.NotNil(pkt) + assert.Equal(uint16(5012), pkt.Header().SequenceNumber) + assert.Equal(uint16(100), queue.Length()) + }) + + t.Run("Updates the length when PopAt* are called", func(t *testing.T) { + queue := NewRingBuffer(128) + pkt := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5000, Timestamp: 500}, Payload: []byte{0x02}} + assert.True(queue.Push(buildRetainablePacket(t, pkt))) + pkt2 := &rtp.Packet{Header: rtp.Header{SequenceNumber: 5001, Timestamp: 501}, Payload: []byte{0x02}} + assert.True(queue.Push(buildRetainablePacket(t, pkt2))) + for i := range 100 { + assert.True(queue.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: uint16(5002 + i), + Timestamp: uint32(502 + i), + }, + Payload: []byte{0x02}, + }))) + } + assert.Equal(uint16(102), queue.Length()) + popped := queue.PopAt(uint16(5002)) + assert.NotNil(popped) + assert.Equal(uint16(5002), popped.Header().SequenceNumber) + assert.Equal(uint16(99), queue.Length()) + + popped = queue.PopAtTimestamp(uint32(504)) + assert.NotNil(popped) + assert.Equal(uint16(5004), popped.Header().SequenceNumber) + assert.Equal(uint16(97), queue.Length()) + }) +} + +func TestRingBuffer_Find(t *testing.T) { + packets := NewRingBuffer(16) + + assert.True(t, packets.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: 1000, + Timestamp: 5, + SSRC: 5, + }, + Payload: []uint8{0xA}, + }))) + + popped := packets.PopAt(1000) + assert.NotNil(t, popped) + + assert.Nil(t, packets.PopAt(1001)) + assert.Nil(t, packets.Peek()) +} + +func TestRingBuffer_Clean(t *testing.T) { + packets := NewRingBuffer(16) + packets.Clear() + assert.True(t, packets.Push(buildRetainablePacket(t, &rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: 1000, + Timestamp: 5, + SSRC: 5, + }, + Payload: []uint8{0xA}, + }))) + assert.EqualValues(t, 1, packets.Length()) + packets.Clear() + assert.EqualValues(t, 0, packets.Length()) + assert.True(t, packets.Empty()) +} + +func TestRingBuffer_Unreference(t *testing.T) { + packets := NewRingBuffer(128) + factory := &rtpbuffer.PacketFactoryNoOp{} + + var refs int64 + finalizer := func(*rtp.Packet) { + atomic.AddInt64(&refs, -1) + } + + numPkts := 100 + for i := range numPkts { + atomic.AddInt64(&refs, 1) + seq := uint16(i) + p := rtp.Packet{ + Header: rtp.Header{ + SequenceNumber: seq, + Timestamp: uint32(i + 42), + }, + Payload: []byte{byte(i)}, + } + runtime.SetFinalizer(&p, finalizer) + rPacket, err := factory.NewPacket(&p.Header, p.Payload, 0, 0) + assert.NoError(t, err) + assert.True(t, packets.Push(rPacket)) + } + for i := 0; i < numPkts-1; i++ { + var popped *rtpbuffer.RetainablePacket + switch i % 3 { + case 0: + popped = packets.Pop() + case 1: + popped = packets.PopAt(uint16(i)) + case 2: + popped = packets.PopAtTimestamp(uint32(i + 42)) + } + if popped != nil { + popped.Release() + } + } + + runtime.GC() + time.Sleep(10 * time.Millisecond) + + remainedRefs := atomic.LoadInt64(&refs) + runtime.KeepAlive(packets) + + // only the last packet should be still referenced + assert.Equal(t, int64(1), remainedRefs) +} + +// Release odd packets and keep even ones, then check if that happened. +func TestRingBuffer_Release_Copy(t *testing.T) { + packets := NewRingBuffer(128) + factory := rtpbuffer.NewPacketFactoryCopy() + + retained := make([]*rtpbuffer.RetainablePacket, 0, 100) + for i := range 100 { + rPacket, err := factory.NewPacket(&rtp.Header{ + SequenceNumber: uint16(i), + Timestamp: uint32(i + 42), + }, []byte{byte(i)}, 0, 0) + assert.NoError(t, err) + assert.True(t, packets.Push(rPacket)) + retained = append(retained, rPacket) + } + + for i := range 100 { + popped := packets.Pop() + assert.NotNil(t, popped) + if i%2 == 1 { + popped.Release() + } + } + + for i, pkt := range retained { + if i%2 == 1 { + assert.Nil(t, pkt.Header()) + assert.Nil(t, pkt.Payload()) + } else { + assert.NotNil(t, pkt.Header()) + assert.NotNil(t, pkt.Payload()) + } + } +} + +func buildRetainablePacket(t *testing.T, pkt *rtp.Packet) *rtpbuffer.RetainablePacket { + t.Helper() + factory := &rtpbuffer.PacketFactoryNoOp{} + rPacket, err := factory.NewPacket(&pkt.Header, pkt.Payload, 0, 0) + assert.NoError(t, err) + + return rPacket +}