From c844d357d59b63baeee4e61765427c66468a9bd7 Mon Sep 17 00:00:00 2001 From: David Chen Date: Sat, 29 Aug 2026 12:33:07 -0700 Subject: [PATCH] Reuse RTP packet storage for FlexFEC recovery --- pkg/sfu/buffer/buffer.go | 23 +- pkg/sfu/buffer/buffer_fec_test.go | 46 ++- pkg/sfu/flexfec/decoder.go | 343 +++++++++------------- pkg/sfu/flexfec/decoder_benchmark_test.go | 30 +- pkg/sfu/flexfec/decoder_test.go | 113 +++++-- 5 files changed, 311 insertions(+), 244 deletions(-) diff --git a/pkg/sfu/buffer/buffer.go b/pkg/sfu/buffer/buffer.go index 9de911a04..cace5961b 100644 --- a/pkg/sfu/buffer/buffer.go +++ b/pkg/sfu/buffer/buffer.go @@ -39,7 +39,8 @@ const ( ) var ( - errInvalidCodec = errors.New("invalid codec") + errInvalidCodec = errors.New("invalid codec") + errFECMediaPacketNotFound = errors.New("fec media packet not found") ) var _ BufferProvider = (*Buffer)(nil) @@ -446,18 +447,30 @@ func (b *Buffer) setFECSSRC(ssrc uint32) { // maybeCreateFECDecoderLocked creates the FEC decoder as soon as the repair // stream SSRC is known and the buffer is bound with a negotiated flexfec -// payload type. Eager creation lets the decoder track media packets before -// the first FEC packet arrives, otherwise the leading protection windows -// would be unrecoverable. +// payload type. Protected media is read from the primary RTP packet bucket. func (b *Buffer) maybeCreateFECDecoderLocked() { if b.fecDecoder != nil || b.fecSSRC == 0 || !b.isBound || b.fecPayloadType == 0 { return } - b.fecDecoder = flexfec.NewDecoder(b.fecSSRC, b.BufferBase.SSRC(), b.logger) + b.fecDecoder = flexfec.NewDecoder(b.fecSSRC, b.BufferBase.SSRC(), b.getFECMediaPacketLocked, b.logger) b.logger.Debugw("flexfec decoder created", "fecSSRC", b.fecSSRC, "mediaSSRC", b.BufferBase.SSRC()) } +func (b *Buffer) getFECMediaPacketLocked(sequenceNumber uint16, dst []byte) (int, error) { + if b.bucket == nil { + return 0, errFECMediaPacketNotFound + } + + headSequenceNumber := b.bucket.HeadSequenceNumber() + extendedSequenceNumber := int64(headSequenceNumber) + int64(int16(sequenceNumber-uint16(headSequenceNumber))) + if extendedSequenceNumber < 0 { + return 0, errFECMediaPacketNotFound + } + + return b.bucket.GetPacket(dst, uint64(extendedSequenceNumber)) +} + // OnFECRecovery is called with counter deltas whenever FEC packets are // processed since the previous callback: FEC packets received, media packets // recovered, FEC packets discarded and FEC bytes received. diff --git a/pkg/sfu/buffer/buffer_fec_test.go b/pkg/sfu/buffer/buffer_fec_test.go index f60c16d67..3d6ebe808 100644 --- a/pkg/sfu/buffer/buffer_fec_test.go +++ b/pkg/sfu/buffer/buffer_fec_test.go @@ -212,9 +212,8 @@ func TestBufferFECRecoveryCallbackCanReenterBuffer(t *testing.T) { func TestBufferFECPairAfterPackets(t *testing.T) { // FEC packets arriving before the ssrc-group is known are queued as - // pending and replayed when the pair is established. Media seen before - // the pairing is not in the decoder window (cold start), so the first - // window is not recoverable, subsequent windows are. + // pending and replayed when the pair is established. Protected media is + // recovered from the primary packet bucket even when pairing is late. factory := NewFactoryOfBufferFactory(500, 200).CreateBufferFactory() primary := factory.GetOrNew(packetio.RTPBufferPacket, fecTestMediaSSRC).(*Buffer) @@ -225,8 +224,9 @@ func TestBufferFECPairAfterPackets(t *testing.T) { fecPackets := encoder.EncodeFec(media, 2) require.NotEmpty(t, fecPackets) + const firstDroppedIdx = 5 for i := range media { - if i == 5 { + if i == firstDroppedIdx { continue } writePacket(t, primary, &media[i]) @@ -243,11 +243,10 @@ func TestBufferFECPairAfterPackets(t *testing.T) { factory.SetFECPair(fecTestFECSSRC, fecTestMediaSSRC) - // pending FEC was replayed into the decoder, no recovery possible for the - // cold-start window + // Pending FEC is replayed and can use media already in the primary bucket. stats = primary.FECDecoderStats() assert.EqualValues(t, len(fecPackets), stats.FECPacketsReceived) - assert.EqualValues(t, 0, stats.PacketsRecovered) + assert.EqualValues(t, 1, stats.PacketsRecovered) // the next window recovers normally media2 := fecTestMediaPackets(t, 210, 10) @@ -266,9 +265,10 @@ func TestBufferFECPairAfterPackets(t *testing.T) { } stats = primary.FECDecoderStats() - assert.EqualValues(t, 1, stats.PacketsRecovered) + assert.EqualValues(t, 2, stats.PacketsRecovered) extSNBySN := readExtSequenceNumbers(t, primary, len(media)-1+len(media2)-1) + requireRecoveredInBucket(t, primary, &media[firstDroppedIdx], extSNBySN, media[0].SequenceNumber) requireRecoveredInBucket(t, primary, &media2[droppedIdx], extSNBySN, media2[0].SequenceNumber) } @@ -304,6 +304,36 @@ func TestBufferFECCoupledBeforeBuffersExist(t *testing.T) { requireRecoveredInBucket(t, primary, &media[droppedIdx], extSNBySN, media[0].SequenceNumber) } +func TestBufferFECSequenceNumberWrap(t *testing.T) { + factory := NewFactoryOfBufferFactory(500, 200).CreateBufferFactory() + factory.SetFECPair(fecTestFECSSRC, fecTestMediaSSRC) + + primary := factory.GetOrNew(packetio.RTPBufferPacket, fecTestMediaSSRC).(*Buffer) + fecBuff := factory.GetOrNew(packetio.RTPBufferPacket, fecTestFECSSRC).(*Buffer) + bindFECTestBuffer(t, primary) + + media := fecTestMediaPackets(t, 65533, 5) + for i := range media { + media[i].Timestamp = 90000 + 3000*uint32(i) + } + fecPackets := pionflexfec.NewFlexEncoder03(fecTestFECPT, fecTestFECSSRC).EncodeFec(media, 1) + require.NotEmpty(t, fecPackets) + + const droppedIdx = 3 // sequence number 0 + for i := range media { + if i != droppedIdx { + writePacket(t, primary, &media[i]) + } + } + for i := range fecPackets { + writePacket(t, fecBuff, &fecPackets[i]) + } + + require.EqualValues(t, 1, primary.FECDecoderStats().PacketsRecovered) + extSNBySN := readExtSequenceNumbers(t, primary, len(media)-1) + requireRecoveredInBucket(t, primary, &media[droppedIdx], extSNBySN, media[0].SequenceNumber) +} + func TestBufferFECIgnoresUnexpectedPayloadType(t *testing.T) { factory := NewFactoryOfBufferFactory(500, 200).CreateBufferFactory() diff --git a/pkg/sfu/flexfec/decoder.go b/pkg/sfu/flexfec/decoder.go index 50ca5b344..b5fe2191b 100644 --- a/pkg/sfu/flexfec/decoder.go +++ b/pkg/sfu/flexfec/decoder.go @@ -20,10 +20,11 @@ // (https://github.com/pion/interceptor, MIT License, Copyright The Pion // community), which is itself modeled on libwebrtc's ForwardErrorCorrection // receiver. Deviations from the pion implementation: -// - packets are deep-copied on insertion (callers reuse packet memory) -// - the media window holds stable heap pointers; the pion version keeps -// values and re-sorts them in place, which invalidates the references -// held by FEC packet state on out-of-order arrival +// - FEC packets are deep-copied only when retained (callers reuse packet +// memory), while protected media is read from the owning packet store +// - packet masks are expanded on demand instead of retaining per-packet +// protection entries +// - recovery XORs the stored RTP wire representation directly // - failed recoveries are not emitted as empty packets // - usage counters for metrics package flexfec @@ -47,17 +48,19 @@ var ( errLastOptionalMaskKBitSetToFalse = errors.New("k-bit of last optional mask is set to false") errEmptyMask = errors.New("empty fec packet mask") errUnknownProtectedSSRC = errors.New("fec is protecting unknown ssrc") + errMediaPacketNotFound = errors.New("protected media packet not found") + errInvalidRecoveredPacketSize = errors.New("invalid recovered packet size") ) const ( - // media window size that triggers the sequence gap reset check - maxMediaPackets = 100 + // number of media arrivals before sequence gaps trigger a state reset + mediaPacketsBeforeGapCheck = 100 // maximum number of FEC packets retained maxFECPackets = 100 - // seen/recovered media packets retained for XOR recovery - recoveredPacketsLimit = 192 // maximum number of sequence numbers represented by the three packet masks maxProtectedPackets = fecMask0Bits + fecMask1Bits + fecMask2Bits + // matches the maximum packet size of the primary RTP packet bucket + maxMediaPacketSize = 1500 ) // FlexFEC-03 header bit fields. @@ -91,6 +94,10 @@ type DecoderStats struct { PacketsRecovered uint64 } +// MediaPacketLookup copies a protected media packet into dst. The caller +// serializes access with writes to the underlying packet store. +type MediaPacketLookup func(sequenceNumber uint16, dst []byte) (int, error) + // Decoder recovers lost media packets of a single protected SSRC from a // FlexFEC-03 repair stream. It is not safe for concurrent use; the owning // buffer serializes access. @@ -98,16 +105,26 @@ type Decoder struct { logger logger.Logger fecSSRC uint32 protectedSSRC uint32 - recoveredPackets []*rtp.Packet + mediaPacketLookup MediaPacketLookup + mediaPacketBuf [maxMediaPacketSize]byte + newestMediaSeq uint16 + mediaPacketsSeen int + hasNewestMediaSeq bool receivedFECPackets []fecPacketState stats DecoderStats } -func NewDecoder(fecSSRC uint32, protectedSSRC uint32, logger logger.Logger) *Decoder { +func NewDecoder( + fecSSRC uint32, + protectedSSRC uint32, + mediaPacketLookup MediaPacketLookup, + logger logger.Logger, +) *Decoder { return &Decoder{ - logger: logger, - fecSSRC: fecSSRC, - protectedSSRC: protectedSSRC, + logger: logger, + fecSSRC: fecSSRC, + protectedSSRC: protectedSSRC, + mediaPacketLookup: mediaPacketLookup, } } @@ -117,50 +134,53 @@ func (d *Decoder) Stats() DecoderStats { // DecodeFec ingests a packet of either the FEC stream (fecSSRC) or the // protected media stream (protectedSSRC) and returns any media packets that -// became recoverable. Returned packets are owned by the decoder's internal -// window; callers must not mutate them. +// became recoverable. Ownership of returned packets transfers to the caller. func (d *Decoder) DecodeFec(receivedPacket *rtp.Packet) []*rtp.Packet { - if receivedPacket.SSRC == d.fecSSRC { + switch receivedPacket.SSRC { + case d.fecSSRC: d.stats.FECPacketsReceived++ d.stats.FECBytesReceived += uint64(len(receivedPacket.Payload)) + d.discardOldFECPackets(receivedPacket.SequenceNumber) + d.insertFECPacket(receivedPacket) + case d.protectedSSRC: + d.observeMediaPacket(receivedPacket.SequenceNumber) + default: + return nil } - // Media packets remain in the recovery window and need an owned copy. FEC - // packets are cloned only if insertFECPacket determines that their state - // must outlive this call. - pkt := receivedPacket - if receivedPacket.SSRC == d.protectedSSRC { - pkt = receivedPacket.Clone() - } - - if len(d.recoveredPackets) >= maxMediaPackets { - backRecoveredPacket := d.recoveredPackets[len(d.recoveredPackets)-1] - if backRecoveredPacket.SSRC == pkt.SSRC { - if seqDiff(pkt.SequenceNumber, backRecoveredPacket.SequenceNumber) > uint16(maxMediaPackets) { - d.logger.Infow("flexfec: big gap in media sequence numbers - resetting buffers") - d.recoveredPackets = nil - d.receivedFECPackets = nil - } - } - } - - d.insertPacket(pkt) - recovered := d.attemptRecovery() d.stats.PacketsRecovered += uint64(len(recovered)) return recovered } -func (d *Decoder) insertPacket(receivedPkt *rtp.Packet) { +func (d *Decoder) observeMediaPacket(sequenceNumber uint16) { + if d.hasNewestMediaSeq && d.mediaPacketsSeen >= mediaPacketsBeforeGapCheck && + seqDiff(sequenceNumber, d.newestMediaSeq) > uint16(mediaPacketsBeforeGapCheck) { + d.logger.Infow("flexfec: big gap in media sequence numbers - resetting buffers") + d.receivedFECPackets = nil + d.mediaPacketsSeen = 0 + d.newestMediaSeq = sequenceNumber + } + + if !d.hasNewestMediaSeq || isNewerSeq(d.newestMediaSeq, sequenceNumber) { + d.newestMediaSeq = sequenceNumber + d.hasNewestMediaSeq = true + } + if d.mediaPacketsSeen < mediaPacketsBeforeGapCheck { + d.mediaPacketsSeen++ + } +} + +func (d *Decoder) discardOldFECPackets(sequenceNumber uint16) { // Discard old FEC packets such that the sequence numbers in // `receivedFECPackets` span at most 1/2 of the sequence number space. // This is important for keeping `receivedFECPackets` sorted, and may // also reduce the possibility of incorrect decoding due to sequence // number wrap-around. - if len(d.receivedFECPackets) > 0 && receivedPkt.SSRC == d.fecSSRC { + if len(d.receivedFECPackets) > 0 { toRemove := 0 for _, fecPkt := range d.receivedFECPackets { - if absInt(int(receivedPkt.SequenceNumber)-int(fecPkt.packet.SequenceNumber)) > 0x3fff { + if absInt(int(sequenceNumber)-int(fecPkt.packet.SequenceNumber)) > 0x3fff { toRemove++ } else { // no need to keep iterating, since receivedFECPackets is sorted @@ -172,47 +192,6 @@ func (d *Decoder) insertPacket(receivedPkt *rtp.Packet) { d.receivedFECPackets = d.receivedFECPackets[toRemove:] } } - - switch receivedPkt.SSRC { - case d.fecSSRC: - d.insertFECPacket(receivedPkt) - case d.protectedSSRC: - d.insertMediaPacket(receivedPkt) - } - - d.discardOldRecoveredPackets() -} - -func (d *Decoder) insertMediaPacket(receivedPkt *rtp.Packet) { - for _, recoveredPacket := range d.recoveredPackets { - if recoveredPacket.SequenceNumber == receivedPkt.SequenceNumber { - return - } - } - - d.recoveredPackets = append(d.recoveredPackets, receivedPkt) - if len(d.recoveredPackets) > 1 && !isNewerSeq( - d.recoveredPackets[len(d.recoveredPackets)-2].SequenceNumber, - receivedPkt.SequenceNumber, - ) { - insertAt := sort.Search(len(d.recoveredPackets)-1, func(i int) bool { - return isNewerSeq(receivedPkt.SequenceNumber, d.recoveredPackets[i].SequenceNumber) - }) - copy(d.recoveredPackets[insertAt+1:], d.recoveredPackets[insertAt:len(d.recoveredPackets)-1]) - d.recoveredPackets[insertAt] = receivedPkt - } - d.updateCoveringFecPackets(receivedPkt) -} - -func (d *Decoder) updateCoveringFecPackets(receivedPkt *rtp.Packet) { - for i := range d.receivedFECPackets { - for j := range d.receivedFECPackets[i].protectedPackets { - pp := &d.receivedFECPackets[i].protectedPackets[j] - if pp.seq == receivedPkt.SequenceNumber { - pp.packet = receivedPkt - } - } - } } func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { @@ -242,13 +221,7 @@ func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { } var protectedSeqBuf [maxProtectedPackets]uint16 - protectedSeqs := appendMaskSequences(protectedSeqBuf[:0], uint64(fec.mask0), fecMask0Bits, fec.seqNumBase) - if fec.mask1 != 0 { - protectedSeqs = appendMaskSequences(protectedSeqs, uint64(fec.mask1), fecMask1Bits, fec.seqNumBase+fecMask0Bits) - } - if fec.mask2 != 0 { - protectedSeqs = appendMaskSequences(protectedSeqs, fec.mask2, fecMask2Bits, fec.seqNumBase+fecMask0Bits+fecMask1Bits) - } + protectedSeqs := fec.protectedSequences(protectedSeqBuf[:0]) if len(protectedSeqs) == 0 { d.stats.FECPacketsDiscarded++ @@ -256,42 +229,10 @@ func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { return } - if countMissingSequences(protectedSeqs, d.recoveredPackets) == 0 { + if d.countMissingPackets(fec, nil) == 0 { return } - protectedPackets := make([]protectedPacket, 0, len(protectedSeqs)) - protectedSeqIt := 0 - recoveredPacketIt := 0 - - for protectedSeqIt < len(protectedSeqs) && recoveredPacketIt < len(d.recoveredPackets) { - switch { - case isNewerSeq(protectedSeqs[protectedSeqIt], d.recoveredPackets[recoveredPacketIt].SequenceNumber): - protectedPackets = append(protectedPackets, protectedPacket{ - seq: protectedSeqs[protectedSeqIt], - packet: nil, - }) - protectedSeqIt++ - case isNewerSeq(d.recoveredPackets[recoveredPacketIt].SequenceNumber, protectedSeqs[protectedSeqIt]): - recoveredPacketIt++ - default: - protectedPackets = append(protectedPackets, protectedPacket{ - seq: protectedSeqs[protectedSeqIt], - packet: d.recoveredPackets[recoveredPacketIt], - }) - protectedSeqIt++ - recoveredPacketIt++ - } - } - - for protectedSeqIt < len(protectedSeqs) { - protectedPackets = append(protectedPackets, protectedPacket{ - seq: protectedSeqs[protectedSeqIt], - packet: nil, - }) - protectedSeqIt++ - } - // The caller may reuse packet memory after DecodeFec returns. Take // ownership only now that this FEC state needs to be retained. ownedFECPkt := fecPkt.Clone() @@ -303,11 +244,7 @@ func (d *Decoder) insertFECPacket(fecPkt *rtp.Packet) { return } - state := fecPacketState{ - packet: ownedFECPkt, - flexFec: ownedFEC, - protectedPackets: protectedPackets, - } + state := fecPacketState{packet: ownedFECPkt, flexFec: ownedFEC} d.receivedFECPackets = append(d.receivedFECPackets, state) if len(d.receivedFECPackets) > 1 && !isNewerSeq( d.receivedFECPackets[len(d.receivedFECPackets)-2].packet.SequenceNumber, @@ -331,7 +268,7 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { packetsRecovered := 0 for i := 0; i < len(d.receivedFECPackets); { fecPkt := &d.receivedFECPackets[i] - packetsMissing := countMissingPackets(fecPkt.protectedPackets) + packetsMissing := d.countMissingPackets(fecPkt.flexFec, recoveredPackets) if packetsMissing == 0 { d.removeFECPacketAt(i) continue @@ -342,7 +279,7 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { continue } - recovered, err := d.recoverPacket(fecPkt) + recovered, err := d.recoverPacket(fecPkt, recoveredPackets) if err != nil { d.logger.Debugw("flexfec: failed to recover packet", "error", err) i++ @@ -351,8 +288,6 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { d.removeFECPacketAt(i) recoveredPackets = append(recoveredPackets, recovered) - d.insertMediaPacket(recovered) - d.discardOldRecoveredPackets() packetsRecovered++ } @@ -364,10 +299,12 @@ func (d *Decoder) attemptRecovery() []*rtp.Packet { return recoveredPackets } -func countMissingPackets(protectedPackets []protectedPacket) int { +func (d *Decoder) countMissingPackets(fec flexFec, recoveredPackets []*rtp.Packet) int { + var protectedSeqBuf [maxProtectedPackets]uint16 + protectedSeqs := fec.protectedSequences(protectedSeqBuf[:0]) missing := 0 - for _, pkt := range protectedPackets { - if pkt.packet == nil { + for _, sequenceNumber := range protectedSeqs { + if _, err := d.getMediaPacket(sequenceNumber, recoveredPackets, d.mediaPacketBuf[:]); err != nil { missing++ if missing > 1 { break @@ -378,26 +315,21 @@ func countMissingPackets(protectedPackets []protectedPacket) int { return missing } -func countMissingSequences(protectedSeqs []uint16, recoveredPackets []*rtp.Packet) int { - missing := 0 - protectedSeqIt := 0 - recoveredPacketIt := 0 - - for protectedSeqIt < len(protectedSeqs) && recoveredPacketIt < len(recoveredPackets) { - switch { - case isNewerSeq(protectedSeqs[protectedSeqIt], recoveredPackets[recoveredPacketIt].SequenceNumber): - missing++ - protectedSeqIt++ - case isNewerSeq(recoveredPackets[recoveredPacketIt].SequenceNumber, protectedSeqs[protectedSeqIt]): - recoveredPacketIt++ - default: - protectedSeqIt++ - recoveredPacketIt++ +func (d *Decoder) getMediaPacket( + sequenceNumber uint16, + recoveredPackets []*rtp.Packet, + dst []byte, +) (int, error) { + for _, recoveredPacket := range recoveredPackets { + if recoveredPacket.SequenceNumber == sequenceNumber { + return recoveredPacket.MarshalTo(dst) } } + if d.mediaPacketLookup == nil { + return 0, errMediaPacketNotFound + } - missing += len(protectedSeqs) - protectedSeqIt - return missing + return d.mediaPacketLookup(sequenceNumber, dst) } func (d *Decoder) removeFECPacketAt(index int) { @@ -407,73 +339,80 @@ func (d *Decoder) removeFECPacketAt(index int) { d.receivedFECPackets = d.receivedFECPackets[:last] } -func (d *Decoder) recoverPacket(fec *fecPacketState) (*rtp.Packet, error) { +func (d *Decoder) recoverPacket(fec *fecPacketState, recoveredPackets []*rtp.Packet) (*rtp.Packet, error) { // https://datatracker.ietf.org/doc/html/draft-ietf-payload-flexible-fec-scheme-03#section-6.3.2 // 2. For the repair packet in T, extract the FEC bit string as the // first 80 bits of the FEC header. - headerRecovery := make([]byte, 12) - copy(headerRecovery, fec.packet.Payload[:10]) + var headerRecovery [12]byte + copy(headerRecovery[:], fec.packet.Payload[:10]) + var protectedSeqBuf [maxProtectedPackets]uint16 + protectedSeqs := fec.flexFec.protectedSequences(protectedSeqBuf[:0]) - var seqnum uint16 - for _, pp := range fec.protectedPackets { - if pp.packet != nil { - // 1. For each of the source packets that are successfully received in - // T, compute the 80-bit string by concatenating the first 64 bits - // of their RTP header and the unsigned network-ordered 16-bit - // representation of their length in bytes minus 12. - receivedHeader, err := pp.packet.Header.Marshal() - if err != nil { - return nil, fmt.Errorf("marshal received header: %w", err) - } - binary.BigEndian.PutUint16(receivedHeader[2:4], uint16(pp.packet.MarshalSize()-12)) - for i := 0; i < 8; i++ { - headerRecovery[i] ^= receivedHeader[i] - } - } else { - seqnum = pp.seq + missing := 0 + var sequenceNumber uint16 + for _, protectedSeq := range protectedSeqs { + n, err := d.getMediaPacket(protectedSeq, recoveredPackets, d.mediaPacketBuf[:]) + if err != nil { + missing++ + sequenceNumber = protectedSeq + continue } + if n < 12 { + return nil, fmt.Errorf("%w: protected packet length %d", errInvalidRecoveredPacketSize, n) + } + + // 1. For each source packet received in T, XOR the first 64 header + // bits with the sequence-number field replaced by the packet length + // after the fixed 12-byte RTP header. + packet := d.mediaPacketBuf[:n] + headerRecovery[0] ^= packet[0] + headerRecovery[1] ^= packet[1] + packetLength := uint16(n - 12) // #nosec G115 -- RTP packet size is bounded above + headerRecovery[2] ^= byte(packetLength >> 8) + headerRecovery[3] ^= byte(packetLength) + for i := 4; i < 8; i++ { + headerRecovery[i] ^= packet[i] + } + } + if missing != 1 { + return nil, fmt.Errorf("cannot recover with %d missing packets", missing) } // set version to 2 headerRecovery[0] |= 0x80 headerRecovery[0] &= 0xbf payloadLength := binary.BigEndian.Uint16(headerRecovery[2:4]) - binary.BigEndian.PutUint16(headerRecovery[2:4], seqnum) + if int(payloadLength)+12 > maxMediaPacketSize { + return nil, fmt.Errorf("%w: recovered packet length %d", errInvalidRecoveredPacketSize, int(payloadLength)+12) + } + binary.BigEndian.PutUint16(headerRecovery[2:4], sequenceNumber) binary.BigEndian.PutUint32(headerRecovery[8:12], d.protectedSSRC) - payloadRecovery := make([]byte, payloadLength) - copy(payloadRecovery, fec.flexFec.payload) - for _, pp := range fec.protectedPackets { - if pp.packet != nil { - packet, err := pp.packet.Marshal() - if err != nil { - return nil, fmt.Errorf("marshal protected packet: %w", err) - } - for i := 0; i < min(int(payloadLength), len(packet)-12); i++ { - payloadRecovery[i] ^= packet[12+i] - } + recoveredRaw := make([]byte, 12+int(payloadLength)) + copy(recoveredRaw[:12], headerRecovery[:]) + copy(recoveredRaw[12:], fec.flexFec.payload) + for _, protectedSeq := range protectedSeqs { + n, err := d.getMediaPacket(protectedSeq, recoveredPackets, d.mediaPacketBuf[:]) + if err != nil { + continue + } + if n < 12 { + return nil, fmt.Errorf("%w: protected packet length %d", errInvalidRecoveredPacketSize, n) + } + for i := 0; i < min(int(payloadLength), n-12); i++ { + recoveredRaw[12+i] ^= d.mediaPacketBuf[12+i] } } - headerRecovery = append(headerRecovery, payloadRecovery...) - packet := &rtp.Packet{} - if err := packet.Unmarshal(headerRecovery); err != nil { + if err := packet.Unmarshal(recoveredRaw); err != nil { return nil, fmt.Errorf("unmarshal recovered: %w", err) } return packet, nil } -func (d *Decoder) discardOldRecoveredPackets() { - if len(d.recoveredPackets) > recoveredPacketsLimit { - toRemove := len(d.recoveredPackets) - recoveredPacketsLimit - clear(d.recoveredPackets[:toRemove]) - d.recoveredPackets = d.recoveredPackets[toRemove:] - } -} - func appendMaskSequences(dst []uint16, mask uint64, bitCount uint16, seqNumBase uint16) []uint16 { for i := uint16(0); i < bitCount; i++ { if (mask>>(bitCount-1-i))&1 == 1 { @@ -485,9 +424,8 @@ func appendMaskSequences(dst []uint16, mask uint64, bitCount uint16, seqNumBase } type fecPacketState struct { - packet *rtp.Packet - flexFec flexFec - protectedPackets []protectedPacket + packet *rtp.Packet + flexFec flexFec } type flexFec struct { @@ -499,9 +437,16 @@ type flexFec struct { payload []byte } -type protectedPacket struct { - seq uint16 - packet *rtp.Packet +func (f flexFec) protectedSequences(dst []uint16) []uint16 { + dst = appendMaskSequences(dst, uint64(f.mask0), fecMask0Bits, f.seqNumBase) + if f.mask1 != 0 { + dst = appendMaskSequences(dst, uint64(f.mask1), fecMask1Bits, f.seqNumBase+fecMask0Bits) + } + if f.mask2 != 0 { + dst = appendMaskSequences(dst, f.mask2, fecMask2Bits, f.seqNumBase+fecMask0Bits+fecMask1Bits) + } + + return dst } func parseFlexFEC03Header(data []byte) (flexFec, error) { diff --git a/pkg/sfu/flexfec/decoder_benchmark_test.go b/pkg/sfu/flexfec/decoder_benchmark_test.go index 6ebf179c8..99a9656fd 100644 --- a/pkg/sfu/flexfec/decoder_benchmark_test.go +++ b/pkg/sfu/flexfec/decoder_benchmark_test.go @@ -40,8 +40,30 @@ func benchmarkMediaPackets(count, payloadSize int) []rtp.Packet { return packets } +func benchmarkMediaLookup(media []rtp.Packet, missingIndex int) MediaPacketLookup { + packets := make(map[uint16][]byte, len(media)) + for i := range media { + if i == missingIndex { + continue + } + raw, err := media[i].Marshal() + if err != nil { + panic(err) + } + packets[media[i].SequenceNumber] = raw + } + + return func(sequenceNumber uint16, dst []byte) (int, error) { + packet, ok := packets[sequenceNumber] + if !ok { + return 0, errTestMediaPacketNotFound + } + return copy(dst, packet), nil + } +} + func BenchmarkDecoderMediaSteadyState1200(b *testing.B) { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := NewDecoder(testFECSSRC, testMediaSSRC, nil, logger.GetLogger()) packet := benchmarkMediaPackets(1, 1200)[0] b.ReportAllocs() @@ -56,11 +78,12 @@ func BenchmarkDecoderMediaSteadyState1200(b *testing.B) { func BenchmarkDecoderCompleteWindow10x1200(b *testing.B) { media := benchmarkMediaPackets(10, 1200) fecPackets := pionflexfec.NewFlexEncoder03(testFECPT, testFECSSRC).EncodeFec(media, 1) + lookup := benchmarkMediaLookup(media, -1) b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := NewDecoder(testFECSSRC, testMediaSSRC, lookup, logger.GetLogger()) for j := range media { decoder.DecodeFec(&media[j]) } @@ -73,11 +96,12 @@ func BenchmarkDecoderCompleteWindow10x1200(b *testing.B) { func BenchmarkDecoderRecoveryWindow10x1200(b *testing.B) { media := benchmarkMediaPackets(10, 1200) fecPackets := pionflexfec.NewFlexEncoder03(testFECPT, testFECSSRC).EncodeFec(media, 1) + lookup := benchmarkMediaLookup(media, 4) b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := NewDecoder(testFECSSRC, testMediaSSRC, lookup, logger.GetLogger()) for j := range media { if j != 4 { decoder.DecodeFec(&media[j]) diff --git a/pkg/sfu/flexfec/decoder_test.go b/pkg/sfu/flexfec/decoder_test.go index 9c55aca4f..22d227b9d 100644 --- a/pkg/sfu/flexfec/decoder_test.go +++ b/pkg/sfu/flexfec/decoder_test.go @@ -16,6 +16,7 @@ package flexfec import ( "encoding/binary" + "errors" "math/rand" "testing" @@ -34,6 +35,51 @@ const ( testMediaPT = uint8(96) ) +var errTestMediaPacketNotFound = errors.New("test media packet not found") + +type testDecoder struct { + *Decoder + mediaPackets map[uint16][]byte +} + +func newTestDecoder(fecSSRC, protectedSSRC uint32, lgr logger.Logger) *testDecoder { + d := &testDecoder{mediaPackets: make(map[uint16][]byte)} + d.Decoder = NewDecoder(fecSSRC, protectedSSRC, d.getMediaPacket, lgr) + return d +} + +func (d *testDecoder) getMediaPacket(sequenceNumber uint16, dst []byte) (int, error) { + packet, ok := d.mediaPackets[sequenceNumber] + if !ok { + return 0, errTestMediaPacketNotFound + } + if len(dst) < len(packet) { + return 0, errors.New("test media packet buffer too small") + } + + return copy(dst, packet), nil +} + +func (d *testDecoder) DecodeFec(packet *rtp.Packet) []*rtp.Packet { + if packet.SSRC == d.protectedSSRC { + raw, err := packet.Marshal() + if err != nil { + panic(err) + } + d.mediaPackets[packet.SequenceNumber] = raw + } + + recovered := d.Decoder.DecodeFec(packet) + for _, recoveredPacket := range recovered { + raw, err := recoveredPacket.Marshal() + if err != nil { + panic(err) + } + d.mediaPackets[recoveredPacket.SequenceNumber] = raw + } + return recovered +} + func makeMediaPackets(t *testing.T, baseSN uint16, count int) []rtp.Packet { t.Helper() rng := rand.New(rand.NewSource(int64(baseSN))) @@ -78,7 +124,7 @@ func TestDecoderRecoversSingleLoss(t *testing.T) { media := makeMediaPackets(t, 100, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) // drop media[2], feed the rest var recovered []*rtp.Packet @@ -104,14 +150,37 @@ func TestDecoderRecoversSingleLoss(t *testing.T) { assert.Equal(t, uint64(0), stats.FECPacketsDiscarded) } +func TestDecoderRecoversPacketWithExtendedHeader(t *testing.T) { + media := makeMediaPackets(t, 150, 5) + for i := range media { + media[i].CSRC = []uint32{uint32(1000 + i)} + require.NoError(t, media[i].SetExtension(3, []byte{byte(i), byte(i + 1)})) + } + fec := encodeFEC(t, media, 1) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + + for i := range media { + if i != 2 { + require.Empty(t, decoder.DecodeFec(&media[i])) + } + } + recovered := decoder.DecodeFec(&fec[0]) + require.Len(t, recovered, 1) + + expectedRaw, err := media[2].Marshal() + require.NoError(t, err) + actualRaw, err := recovered[0].Marshal() + require.NoError(t, err) + assert.Equal(t, expectedRaw, actualRaw) +} + func TestDecoderRecoversWithLateMedia(t *testing.T) { // FEC arrives while two packets are missing; recovery happens once one of - // them shows up late. Exercises updateCoveringFecPackets and the - // stable-pointer window. + // them shows up late. Exercises retained FEC state and packet lookup. media := makeMediaPackets(t, 200, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) var recovered []*rtp.Packet for _, i := range []int{0, 3, 4} { @@ -132,7 +201,7 @@ func TestDecoderRecoversWithLateMedia(t *testing.T) { } func TestDecoderRecoversMultipleWindows(t *testing.T) { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) encoder := pionflexfec.NewFlexEncoder03(testFECPT, testFECSSRC) var allRecovered []*rtp.Packet @@ -170,7 +239,7 @@ func TestDecoderSequenceNumberWrap(t *testing.T) { media := makeMediaPackets(t, 65533, 5) // spans 65533..1 fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) var recovered []*rtp.Packet for i := range media { @@ -187,24 +256,10 @@ func TestDecoderSequenceNumberWrap(t *testing.T) { requirePacketEqual(t, &media[3], recovered[0]) } -func TestDecoderMediaWindowOrder(t *testing.T) { - media := makeMediaPackets(t, 65534, 4) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) - - for _, i := range []int{1, 3, 0, 2} { - require.Empty(t, decoder.DecodeFec(&media[i])) - } - - require.Len(t, decoder.recoveredPackets, 4) - for i := range media { - assert.Equal(t, media[i].SequenceNumber, decoder.recoveredPackets[i].SequenceNumber) - } -} - func TestDecoderFECWindowOrder(t *testing.T) { media := makeMediaPackets(t, 50, 5) fec := encodeFEC(t, media, 1)[0] - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for _, i := range []int{0, 3, 4} { require.Empty(t, decoder.DecodeFec(&media[i])) @@ -225,7 +280,7 @@ func TestDecoderDiscardsForeignProtectedSSRC(t *testing.T) { fec := encodeFEC(t, media, 1) // decoder bound to a different protected stream - decoder := NewDecoder(testFECSSRC, testMediaSSRC+1, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC+1, logger.GetLogger()) recovered := decoder.DecodeFec(&fec[0]) require.Empty(t, recovered) @@ -235,7 +290,7 @@ func TestDecoderDiscardsForeignProtectedSSRC(t *testing.T) { } func TestDecoderDiscardsMalformedFEC(t *testing.T) { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for _, payload := range [][]byte{ nil, @@ -283,7 +338,7 @@ func TestDecoderDiscardsDuplicateFEC(t *testing.T) { media := makeMediaPackets(t, 400, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for _, i := range []int{0, 3, 4} { decoder.DecodeFec(&media[i]) } @@ -299,7 +354,7 @@ func TestDecoderDoesNotRetainCompleteFECState(t *testing.T) { media := makeMediaPackets(t, 450, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) for i := range media { require.Empty(t, decoder.DecodeFec(&media[i])) } @@ -313,7 +368,7 @@ func TestDecoderInputMemoryReuse(t *testing.T) { media := makeMediaPackets(t, 500, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) scratch := &rtp.Packet{} feed := func(src *rtp.Packet) []*rtp.Packet { @@ -347,7 +402,7 @@ func TestDecoderRetainedFECMemoryReuse(t *testing.T) { media := makeMediaPackets(t, 550, 5) fec := encodeFEC(t, media, 1) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) scratch := &rtp.Packet{} feed := func(src *rtp.Packet) []*rtp.Packet { buf, err := src.Marshal() @@ -378,7 +433,7 @@ func TestDecoderTwoFECPacketsTwoLosses(t *testing.T) { fec := encodeFEC(t, media, 2) require.Len(t, fec, 2) - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) var recovered []*rtp.Packet for i := range media { @@ -406,7 +461,7 @@ func TestDecoderTwoFECPacketsTwoLosses(t *testing.T) { } func TestDecoderResetsOnBigSequenceGap(t *testing.T) { - decoder := NewDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) + decoder := newTestDecoder(testFECSSRC, testMediaSSRC, logger.GetLogger()) media := makeMediaPackets(t, 100, 110) for i := range media {