From b0c39fca22287010661de201386eb92ff14d0701 Mon Sep 17 00:00:00 2001 From: Juliusz Chroboczek Date: Sun, 11 Jul 2021 23:09:19 +0200 Subject: [PATCH] Implement loss handling in diskwriter. --- diskwriter/diskwriter.go | 63 ++++++++++++++++++++++++++++++++-------- 1 file changed, 51 insertions(+), 12 deletions(-) diff --git a/diskwriter/diskwriter.go b/diskwriter/diskwriter.go index 1f88904..bff4449 100644 --- a/diskwriter/diskwriter.go +++ b/diskwriter/diskwriter.go @@ -255,7 +255,7 @@ type maybeUint32 uint64 const none maybeUint32 = 0 func some(value uint32) maybeUint32 { - return maybeUint32(uint64(1 << 32) | uint64(value)) + return maybeUint32(uint64(1<<32) | uint64(value)) } func valid(m maybeUint32) bool { @@ -270,10 +270,10 @@ type diskTrack struct { remote conn.UpTrack conn *diskConn - writer webm.BlockWriteCloser - builder *samplebuilder.SampleBuilder - - origin maybeUint32 + writer webm.BlockWriteCloser + builder *samplebuilder.SampleBuilder + lastSeqno maybeUint32 + origin maybeUint32 kfRequested time.Time lastKf time.Time @@ -476,8 +476,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { return 0, nil } - codec := t.remote.Codec() - + // samplebuilder retains packets data := make([]byte, len(buf)) copy(data, buf) p := new(rtp.Packet) @@ -487,6 +486,46 @@ func (t *diskTrack) Write(buf []byte) (int, error) { return 0, nil } + if valid(t.lastSeqno) { + lastSeqno := uint16(value(t.lastSeqno)) + if ((p.SequenceNumber - lastSeqno) & 0x8000) == 0 { + count := p.SequenceNumber - lastSeqno + if count > 0 && count < 128 { + var nacks []uint16 + for i := lastSeqno + 1; i != p.SequenceNumber; i++ { + // different buf each time + buf := make([]byte, 1504) + n := t.remote.GetRTP(i, buf) + if n > 0 { + p := new(rtp.Packet) + err := p.Unmarshal(buf) + if err == nil { + t.writeRTP(p) + } + } else { + nacks = append(nacks, i) + } + } + if len(nacks) > 0 { + t.remote.Nack(t.conn.remote, nacks) + } + } + } + } + + t.lastSeqno = some(uint32(p.SequenceNumber)) + + err = t.writeRTP(p) + if err != nil { + return 0, err + } + return len(buf), nil +} + +// writeRTP writes the packet without doing any loss recovery. +// Called locked. +func (t *diskTrack) writeRTP(p *rtp.Packet) error { + codec := t.remote.Codec() if strings.EqualFold(codec.MimeType, "video/vp9") { var vp9 codecs.VP9Packet _, err := vp9.Unmarshal(p.Payload) @@ -509,7 +548,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { for { sample, ts := t.builder.PopWithTimestamp() if sample == nil { - return len(buf), nil + return nil } keyframe := true @@ -528,7 +567,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { t.conn.warn( "Write to disk " + err.Error(), ) - return 0, err + return err } } } else { @@ -540,7 +579,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { "Write to disk " + err.Error(), ) - return 0, err + return err } } } @@ -554,7 +593,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { t.remote.RequestKeyframe() t.kfRequested = now } - return 0, nil + return nil } if t.writer == nil { @@ -569,7 +608,7 @@ func (t *diskTrack) Write(buf []byte) (int, error) { tm := ts / (t.remote.Codec().ClockRate / 1000) _, err := t.writer.Write(keyframe, int64(tm), sample.Data) if err != nil { - return 0, err + return err } } }