diff --git a/go.mod b/go.mod index 7fbbd32c..73dcae19 100644 --- a/go.mod +++ b/go.mod @@ -4,7 +4,7 @@ go 1.17 require ( code.cloudfoundry.org/bytefmt v0.0.0-20211005130812-5bb3c17173e5 - github.com/aler9/gortsplib v0.0.0-20220401091943-cec5326ccfed + github.com/aler9/gortsplib v0.0.0-20220408160915-2d2e62f55bae github.com/asticode/go-astits v1.10.0 github.com/fsnotify/fsnotify v1.4.9 github.com/gin-gonic/gin v1.7.2 diff --git a/go.sum b/go.sum index 6a420fd2..2e8dee2e 100644 --- a/go.sum +++ b/go.sum @@ -4,8 +4,8 @@ github.com/alecthomas/template v0.0.0-20190718012654-fb15b899a751 h1:JYp7IbQjafo github.com/alecthomas/template v0.0.0-20190718012654-fb15b899a751/go.mod h1:LOuyumcjzFXgccqObfd/Ljyb9UuFJ6TxHnclSeseNhc= github.com/alecthomas/units v0.0.0-20190924025748-f65c72e2690d h1:UQZhZ2O0vMHr2cI+DC1Mbh0TJxzA3RcLoMsFw+aXw7E= github.com/alecthomas/units v0.0.0-20190924025748-f65c72e2690d/go.mod h1:rBZYJk541a8SKzHPHnH3zbiI+7dagKZ0cgpgrD7Fyho= -github.com/aler9/gortsplib v0.0.0-20220401091943-cec5326ccfed h1:lA/dMUwmQcBCeuFRYkPr3qnGPKgrFdGC0RZzB8tRKXw= -github.com/aler9/gortsplib v0.0.0-20220401091943-cec5326ccfed/go.mod h1:4mWq8mM6v8KrSQG4sEdnvM6+ZVKPPgKtf75TYR+jsKQ= +github.com/aler9/gortsplib v0.0.0-20220408160915-2d2e62f55bae h1:BGe90r+y1BRvSz1b1OIbee0q9c2MdI2GUhnzVm0XoSU= +github.com/aler9/gortsplib v0.0.0-20220408160915-2d2e62f55bae/go.mod h1:Mezkz7Jb5zrIWP6MxJ2uBgt5xwywZkcdmuQZ2QrFYsM= github.com/aler9/rtmp v0.0.0-20210403095203-3be4a5535927 h1:95mXJ5fUCYpBRdSOnLAQAdJHHKxxxJrVCiaqDi965YQ= github.com/aler9/rtmp v0.0.0-20210403095203-3be4a5535927/go.mod h1:vzuE21rowz+lT1NGsWbreIvYulgBpCGnQyeTyFblUHc= github.com/asticode/go-astikit v0.20.0 h1:+7N+J4E4lWx2QOkRdOf6DafWJMv6O4RRfgClwQokrH8= diff --git a/internal/core/data.go b/internal/core/data.go new file mode 100644 index 00000000..a6b30969 --- /dev/null +++ b/internal/core/data.go @@ -0,0 +1,14 @@ +package core + +import ( + "time" + + "github.com/pion/rtp" +) + +type data struct { + rtp *rtp.Packet + ptsEqualsDTS bool + h264NALUs [][]byte + h264PTS time.Duration +} diff --git a/internal/core/hls_muxer.go b/internal/core/hls_muxer.go index d880b097..5dec922f 100644 --- a/internal/core/hls_muxer.go +++ b/internal/core/hls_muxer.go @@ -16,8 +16,6 @@ import ( "github.com/aler9/gortsplib" "github.com/aler9/gortsplib/pkg/ringbuffer" "github.com/aler9/gortsplib/pkg/rtpaac" - "github.com/aler9/gortsplib/pkg/rtph264" - "github.com/pion/rtp" "github.com/aler9/rtsp-simple-server/internal/conf" "github.com/aler9/rtsp-simple-server/internal/hls" @@ -107,9 +105,9 @@ type hlsMuxerRequest struct { res chan hlsMuxerResponse } -type hlsMuxerTrackIDPayloadPair struct { +type hlsMuxerTrackIDDataPair struct { trackID int - packet *rtp.Packet + data *data } type hlsMuxerPathManager interface { @@ -282,7 +280,6 @@ func (m *hlsMuxer) runInner(innerCtx context.Context, innerReady chan struct{}) var videoTrack *gortsplib.TrackH264 videoTrackID := -1 - var h264Decoder *rtph264.Decoder var audioTrack *gortsplib.TrackAAC audioTrackID := -1 var aacDecoder *rtpaac.Decoder @@ -296,8 +293,6 @@ func (m *hlsMuxer) runInner(innerCtx context.Context, innerReady chan struct{}) videoTrack = tt videoTrackID = i - h264Decoder = &rtph264.Decoder{} - h264Decoder.Init() case *gortsplib.TrackAAC: if audioTrack != nil { @@ -342,25 +337,20 @@ func (m *hlsMuxer) runInner(innerCtx context.Context, innerReady chan struct{}) if !ok { return fmt.Errorf("terminated") } - pair := data.(hlsMuxerTrackIDPayloadPair) + pair := data.(hlsMuxerTrackIDDataPair) if videoTrack != nil && pair.trackID == videoTrackID { - nalus, pts, err := h264Decoder.DecodeUntilMarker(pair.packet) - if err != nil { - if err != rtph264.ErrMorePacketsNeeded && - err != rtph264.ErrNonStartingPacketAndNoPrevious { - m.log(logger.Warn, "unable to decode video track: %v", err) - } + if pair.data.h264NALUs == nil { continue } - err = m.muxer.WriteH264(pts, nalus) + err = m.muxer.WriteH264(pair.data.h264PTS, pair.data.h264NALUs) if err != nil { m.log(logger.Warn, "unable to write segment: %v", err) continue } } else if audioTrack != nil && pair.trackID == audioTrackID { - aus, pts, err := aacDecoder.Decode(pair.packet) + aus, pts, err := aacDecoder.Decode(pair.data.rtp) if err != nil { if err != rtpaac.ErrMorePacketsNeeded { m.log(logger.Warn, "unable to decode audio track: %v", err) @@ -536,9 +526,9 @@ func (m *hlsMuxer) onReaderAccepted() { m.log(logger.Info, "is converting into HLS") } -// onReaderPacketRTP implements reader. -func (m *hlsMuxer) onReaderPacketRTP(trackID int, pkt *rtp.Packet) { - m.ringBuffer.Push(hlsMuxerTrackIDPayloadPair{trackID, pkt}) +// onReaderData implements reader. +func (m *hlsMuxer) onReaderData(trackID int, data *data) { + m.ringBuffer.Push(hlsMuxerTrackIDDataPair{trackID, data}) } // onReaderAPIDescribe implements reader. diff --git a/internal/core/hls_source.go b/internal/core/hls_source.go index 64455035..a6122cfc 100644 --- a/internal/core/hls_source.go +++ b/internal/core/hls_source.go @@ -6,6 +6,7 @@ import ( "time" "github.com/aler9/gortsplib" + "github.com/aler9/gortsplib/pkg/h264" "github.com/aler9/gortsplib/pkg/rtpaac" "github.com/aler9/gortsplib/pkg/rtph264" @@ -146,8 +147,21 @@ func (s *hlsSource) runInner() bool { return } - for _, pkt := range pkts { - stream.writePacketRTP(videoTrackID, pkt) + lastPkt := len(pkts) - 1 + for i, pkt := range pkts { + if i != lastPkt { + stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: false, + }) + } else { + stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: h264.IDRPresent(nalus), + h264NALUs: nalus, + h264PTS: pts, + }) + } } } @@ -162,7 +176,10 @@ func (s *hlsSource) runInner() bool { } for _, pkt := range pkts { - stream.writePacketRTP(audioTrackID, pkt) + stream.writeData(audioTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: true, + }) } } diff --git a/internal/core/hls_source_test.go b/internal/core/hls_source_test.go index 18bff390..ba1caced 100644 --- a/internal/core/hls_source_test.go +++ b/internal/core/hls_source_test.go @@ -13,7 +13,6 @@ import ( "github.com/aler9/gortsplib/pkg/h264" "github.com/asticode/go-astits" "github.com/gin-gonic/gin" - "github.com/pion/rtp" "github.com/stretchr/testify/require" ) @@ -134,8 +133,8 @@ func TestHLSSource(t *testing.T) { frameRecv := make(chan struct{}) c := gortsplib.Client{ - OnPacketRTP: func(trackID int, pkt *rtp.Packet) { - require.Equal(t, []byte{0x05}, pkt.Payload) + OnPacketRTP: func(ctx *gortsplib.ClientOnPacketRTPCtx) { + require.Equal(t, []byte{0x05}, ctx.Packet.Payload) close(frameRecv) }, } diff --git a/internal/core/reader.go b/internal/core/reader.go index 2f632aff..e7aba880 100644 --- a/internal/core/reader.go +++ b/internal/core/reader.go @@ -1,13 +1,9 @@ package core -import ( - "github.com/pion/rtp" -) - // reader is an entity that can read a stream. type reader interface { close() onReaderAccepted() - onReaderPacketRTP(int, *rtp.Packet) + onReaderData(int, *data) onReaderAPIDescribe() interface{} } diff --git a/internal/core/rtmp_conn.go b/internal/core/rtmp_conn.go index c65742da..6e91a52b 100644 --- a/internal/core/rtmp_conn.go +++ b/internal/core/rtmp_conn.go @@ -16,7 +16,6 @@ import ( "github.com/aler9/gortsplib/pkg/rtpaac" "github.com/aler9/gortsplib/pkg/rtph264" "github.com/notedit/rtmp/av" - "github.com/pion/rtp" "github.com/aler9/rtsp-simple-server/internal/conf" "github.com/aler9/rtsp-simple-server/internal/externalcmd" @@ -44,9 +43,9 @@ const ( rtmpConnStatePublish ) -type rtmpConnTrackIDPayloadPair struct { +type rtmpConnTrackIDDataPair struct { trackID int - packet *rtp.Packet + data *data } type rtmpConnPathManager interface { @@ -260,7 +259,6 @@ func (c *rtmpConn) runRead(ctx context.Context) error { var videoTrack *gortsplib.TrackH264 videoTrackID := -1 - var h264Decoder *rtph264.Decoder var audioTrack *gortsplib.TrackAAC audioTrackID := -1 var aacDecoder *rtpaac.Decoder @@ -274,8 +272,6 @@ func (c *rtmpConn) runRead(ctx context.Context) error { videoTrack = tt videoTrackID = i - h264Decoder = &rtph264.Decoder{} - h264Decoder.Init() case *gortsplib.TrackAAC: if audioTrack != nil { @@ -338,20 +334,16 @@ func (c *rtmpConn) runRead(ctx context.Context) error { if !ok { return fmt.Errorf("terminated") } - pair := data.(rtmpConnTrackIDPayloadPair) + pair := data.(rtmpConnTrackIDDataPair) if videoTrack != nil && pair.trackID == videoTrackID { - nalus, pts, err := h264Decoder.DecodeUntilMarker(pair.packet) - if err != nil { - if err != rtph264.ErrMorePacketsNeeded && err != rtph264.ErrNonStartingPacketAndNoPrevious { - c.log(logger.Warn, "unable to decode video track: %v", err) - } + if pair.data.h264NALUs == nil { continue } var nalusFiltered [][]byte - for _, nalu := range nalus { + for _, nalu := range pair.data.h264NALUs { // remove SPS, PPS and AUD, not needed by RTMP typ := h264.NALUType(nalu[0] & 0x1F) switch typ { @@ -362,24 +354,14 @@ func (c *rtmpConn) runRead(ctx context.Context) error { nalusFiltered = append(nalusFiltered, nalu) } - idrPresent := func() bool { - for _, nalu := range nalus { - typ := h264.NALUType(nalu[0] & 0x1F) - if typ == h264.NALUTypeIDR { - return true - } - } - return false - }() - // wait until we receive an IDR if !videoFirstIDRFound { - if !idrPresent { + if !h264.IDRPresent(nalusFiltered) { continue } videoFirstIDRFound = true - videoStartPTS = pts + videoStartPTS = pair.data.h264PTS videoDTSEst = h264.NewDTSEstimator() } @@ -388,7 +370,7 @@ func (c *rtmpConn) runRead(ctx context.Context) error { return err } - pts -= videoStartPTS + pts := pair.data.h264PTS - videoStartPTS dts := videoDTSEst.Feed(pts) c.conn.SetWriteDeadline(time.Now().Add(time.Duration(c.writeTimeout))) @@ -402,7 +384,7 @@ func (c *rtmpConn) runRead(ctx context.Context) error { return err } } else if audioTrack != nil && pair.trackID == audioTrackID { - aus, pts, err := aacDecoder.Decode(pair.packet) + aus, pts, err := aacDecoder.Decode(pair.data.rtp) if err != nil { if err != rtpaac.ErrMorePacketsNeeded { c.log(logger.Warn, "unable to decode audio track: %v", err) @@ -545,13 +527,28 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { continue } - pkts, err := h264Encoder.Encode(outNALUs, pkt.Time+pkt.CTime) + pts := pkt.Time + pkt.CTime + + pkts, err := h264Encoder.Encode(outNALUs, pts) if err != nil { return fmt.Errorf("error while encoding H264: %v", err) } - for _, pkt := range pkts { - rres.stream.writePacketRTP(videoTrackID, pkt) + lastPkt := len(pkts) - 1 + for i, pkt := range pkts { + if i != lastPkt { + rres.stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: false, + }) + } else { + rres.stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: h264.IDRPresent(outNALUs), + h264NALUs: outNALUs, + h264PTS: pts, + }) + } } case av.AAC: @@ -565,7 +562,10 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { } for _, pkt := range pkts { - rres.stream.writePacketRTP(audioTrackID, pkt) + rres.stream.writeData(audioTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: true, + }) } } } @@ -622,9 +622,9 @@ func (c *rtmpConn) onReaderAccepted() { c.log(logger.Info, "is reading from path '%s'", c.path.Name()) } -// onReaderPacketRTP implements reader. -func (c *rtmpConn) onReaderPacketRTP(trackID int, pkt *rtp.Packet) { - c.ringBuffer.Push(rtmpConnTrackIDPayloadPair{trackID, pkt}) +// onReaderData implements reader. +func (c *rtmpConn) onReaderData(trackID int, data *data) { + c.ringBuffer.Push(rtmpConnTrackIDDataPair{trackID, data}) } // onReaderAPIDescribe implements reader. diff --git a/internal/core/rtmp_source.go b/internal/core/rtmp_source.go index 95d726f9..6e3bccf5 100644 --- a/internal/core/rtmp_source.go +++ b/internal/core/rtmp_source.go @@ -165,6 +165,7 @@ func (s *rtmpSource) runInner() bool { defer func() { s.parent.onSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{source: s}) }() + for { conn.SetReadDeadline(time.Now().Add(time.Duration(s.readTimeout))) pkt, err := conn.ReadPacket() @@ -195,13 +196,28 @@ func (s *rtmpSource) runInner() bool { outNALUs = append(outNALUs, nalu) } - pkts, err := h264Encoder.Encode(outNALUs, pkt.Time+pkt.CTime) + pts := pkt.Time + pkt.CTime + + pkts, err := h264Encoder.Encode(outNALUs, pts) if err != nil { return fmt.Errorf("error while encoding H264: %v", err) } - for _, pkt := range pkts { - res.stream.writePacketRTP(videoTrackID, pkt) + lastPkt := len(pkts) - 1 + for i, pkt := range pkts { + if i != lastPkt { + res.stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: false, + }) + } else { + res.stream.writeData(videoTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: h264.IDRPresent(outNALUs), + h264NALUs: outNALUs, + h264PTS: pts, + }) + } } case av.AAC: @@ -215,7 +231,10 @@ func (s *rtmpSource) runInner() bool { } for _, pkt := range pkts { - res.stream.writePacketRTP(audioTrackID, pkt) + res.stream.writeData(audioTrackID, &data{ + rtp: pkt, + ptsEqualsDTS: true, + }) } } } diff --git a/internal/core/rtsp_server_test.go b/internal/core/rtsp_server_test.go index d899dd6a..ef6a528f 100644 --- a/internal/core/rtsp_server_test.go +++ b/internal/core/rtsp_server_test.go @@ -459,11 +459,11 @@ func TestRTSPServerPublisherOverride(t *testing.T) { frameRecv := make(chan struct{}) c := gortsplib.Client{ - OnPacketRTP: func(trackID int, pkt *rtp.Packet) { + OnPacketRTP: func(ctx *gortsplib.ClientOnPacketRTPCtx) { if ca == "enabled" { - require.Equal(t, []byte{0x05, 0x06, 0x07, 0x08}, pkt.Payload) + require.Equal(t, []byte{0x05, 0x06, 0x07, 0x08}, ctx.Packet.Payload) } else { - require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, pkt.Payload) + require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, ctx.Packet.Payload) } close(frameRecv) }, @@ -483,7 +483,7 @@ func TestRTSPServerPublisherOverride(t *testing.T) { Marker: true, }, Payload: []byte{0x01, 0x02, 0x03, 0x04}, - }) + }, true) if ca == "enabled" { require.Error(t, err) } else { @@ -501,7 +501,7 @@ func TestRTSPServerPublisherOverride(t *testing.T) { Marker: true, }, Payload: []byte{0x05, 0x06, 0x07, 0x08}, - }) + }, true) require.NoError(t, err) } diff --git a/internal/core/rtsp_session.go b/internal/core/rtsp_session.go index 0e38c3bf..5c597f07 100644 --- a/internal/core/rtsp_session.go +++ b/internal/core/rtsp_session.go @@ -9,7 +9,6 @@ import ( "github.com/aler9/gortsplib" "github.com/aler9/gortsplib/pkg/base" - "github.com/pion/rtp" "github.com/aler9/rtsp-simple-server/internal/conf" "github.com/aler9/rtsp-simple-server/internal/externalcmd" @@ -42,10 +41,9 @@ type rtspSession struct { path *path state gortsplib.ServerSessionState stateMutex sync.Mutex - setuppedTracks map[int]gortsplib.Track // read - onReadCmd *externalcmd.Cmd // read - announcedTracks gortsplib.Tracks // publish - stream *stream // publish + onReadCmd *externalcmd.Cmd // read + announcedTracks gortsplib.Tracks // publish + stream *stream // publish } func newRTSPSession( @@ -231,11 +229,6 @@ func (s *rtspSession) onSetup(c *rtspConn, ctx *gortsplib.ServerHandlerOnSetupCt }, nil, fmt.Errorf("track %d does not exist", ctx.TrackID) } - if s.setuppedTracks == nil { - s.setuppedTracks = make(map[int]gortsplib.Track) - } - s.setuppedTracks[ctx.TrackID] = res.stream.tracks()[ctx.TrackID] - s.stateMutex.Lock() s.state = gortsplib.ServerSessionStatePrePlay s.stateMutex.Unlock() @@ -348,8 +341,8 @@ func (s *rtspSession) onReaderAccepted() { s.ss.SetuppedTransport()) } -// onReaderPacketRTP implements reader. -func (s *rtspSession) onReaderPacketRTP(trackID int, pkt *rtp.Packet) { +// onReaderData implements reader. +func (s *rtspSession) onReaderData(trackID int, data *data) { // packets are routed to the session by gortsplib.ServerStream. } @@ -399,5 +392,17 @@ func (s *rtspSession) onPublisherAccepted(tracksLen int) { // onPacketRTP is called by rtspServer. func (s *rtspSession) onPacketRTP(ctx *gortsplib.ServerHandlerOnPacketRTPCtx) { - s.stream.writePacketRTP(ctx.TrackID, ctx.Packet) + if ctx.H264NALUs != nil { + s.stream.writeData(ctx.TrackID, &data{ + rtp: ctx.Packet, + ptsEqualsDTS: ctx.PTSEqualsDTS, + h264NALUs: append([][]byte(nil), ctx.H264NALUs...), + h264PTS: ctx.H264PTS, + }) + } else { + s.stream.writeData(ctx.TrackID, &data{ + rtp: ctx.Packet, + ptsEqualsDTS: ctx.PTSEqualsDTS, + }) + } } diff --git a/internal/core/rtsp_source.go b/internal/core/rtsp_source.go index 4d6ab23e..a60f3896 100644 --- a/internal/core/rtsp_source.go +++ b/internal/core/rtsp_source.go @@ -12,9 +12,6 @@ import ( "github.com/aler9/gortsplib" "github.com/aler9/gortsplib/pkg/base" - "github.com/aler9/gortsplib/pkg/h264" - "github.com/aler9/gortsplib/pkg/rtph264" - "github.com/pion/rtp" "github.com/aler9/rtsp-simple-server/internal/conf" "github.com/aler9/rtsp-simple-server/internal/logger" @@ -186,11 +183,6 @@ func (s *rtspSource) runInner() bool { } } - err = s.handleMissingH264Params(c, tracks) - if err != nil { - return err - } - res := s.parent.onSourceStaticSetReady(pathSourceStaticSetReadyReq{ source: s, tracks: c.Tracks(), @@ -205,8 +197,20 @@ func (s *rtspSource) runInner() bool { s.parent.onSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{source: s}) }() - c.OnPacketRTP = func(trackID int, pkt *rtp.Packet) { - res.stream.writePacketRTP(trackID, pkt) + c.OnPacketRTP = func(ctx *gortsplib.ClientOnPacketRTPCtx) { + if ctx.H264NALUs != nil { + res.stream.writeData(ctx.TrackID, &data{ + rtp: ctx.Packet, + ptsEqualsDTS: ctx.PTSEqualsDTS, + h264NALUs: append([][]byte(nil), ctx.H264NALUs...), + h264PTS: ctx.H264PTS, + }) + } else { + res.stream.writeData(ctx.TrackID, &data{ + rtp: ctx.Packet, + ptsEqualsDTS: ctx.PTSEqualsDTS, + }) + } } _, err = c.Play(nil) @@ -230,131 +234,6 @@ func (s *rtspSource) runInner() bool { } } -func (s *rtspSource) handleMissingH264Params(c *gortsplib.Client, tracks gortsplib.Tracks) error { - h264Track, h264TrackID := func() (*gortsplib.TrackH264, int) { - for i, t := range tracks { - if th264, ok := t.(*gortsplib.TrackH264); ok { - if th264.SPS() == nil { - return th264, i - } - } - } - return nil, -1 - }() - if h264TrackID < 0 { - return nil - } - - if h264Track.SPS() != nil && h264Track.PPS() != nil { - return nil - } - - s.log(logger.Info, "source has not provided H264 parameters (SPS and PPS)"+ - " inside the SDP; extracting them from the stream...") - - var streamMutex sync.RWMutex - var stream *stream - decoder := &rtph264.Decoder{} - decoder.Init() - var sps []byte - var pps []byte - paramsReceived := make(chan struct{}) - - c.OnPacketRTP = func(trackID int, pkt *rtp.Packet) { - streamMutex.RLock() - defer streamMutex.RUnlock() - - if stream == nil { - if trackID != h264TrackID { - return - } - - select { - case <-paramsReceived: - return - default: - } - - nalus, _, err := decoder.Decode(pkt) - if err != nil { - return - } - - for _, nalu := range nalus { - typ := h264.NALUType(nalu[0] & 0x1F) - switch typ { - case h264.NALUTypeSPS: - sps = nalu - if sps != nil && pps != nil { - close(paramsReceived) - } - - case h264.NALUTypePPS: - pps = nalu - if sps != nil && pps != nil { - close(paramsReceived) - } - } - } - } else { - stream.writePacketRTP(trackID, pkt) - } - } - - _, err := c.Play(nil) - if err != nil { - return err - } - - readErr := make(chan error) - go func() { - readErr <- c.Wait() - }() - - timeout := time.NewTimer(15 * time.Second) - defer timeout.Stop() - - select { - case err := <-readErr: - return err - - case <-timeout.C: - c.Close() - <-readErr - return fmt.Errorf("source did not send H264 parameters in time") - - case <-paramsReceived: - s.log(logger.Info, "H264 parameters extracted") - - h264Track.SetSPS(sps) - h264Track.SetPPS(pps) - - res := s.parent.onSourceStaticSetReady(pathSourceStaticSetReadyReq{ - source: s, - tracks: tracks, - }) - if res.err != nil { - c.Close() - <-readErr - return res.err - } - - func() { - streamMutex.Lock() - defer streamMutex.Unlock() - stream = res.stream - }() - - s.log(logger.Info, "ready") - - defer func() { - s.parent.onSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{source: s}) - }() - - return <-readErr - } -} - // onSourceAPIDescribe implements source. func (*rtspSource) onSourceAPIDescribe() interface{} { return struct { diff --git a/internal/core/rtsp_source_test.go b/internal/core/rtsp_source_test.go index dd590939..12a64c5f 100644 --- a/internal/core/rtsp_source_test.go +++ b/internal/core/rtsp_source_test.go @@ -84,7 +84,7 @@ func TestRTSPSource(t *testing.T) { Marker: true, }, Payload: []byte{0x01, 0x02, 0x03, 0x04}, - }) + }, true) }() return &base.Response{ @@ -143,8 +143,8 @@ func TestRTSPSource(t *testing.T) { received := make(chan struct{}) c := gortsplib.Client{ - OnPacketRTP: func(trackID int, pkt *rtp.Packet) { - require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, pkt.Payload) + OnPacketRTP: func(ctx *gortsplib.ClientOnPacketRTPCtx) { + require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, ctx.Packet.Payload) close(received) }, } @@ -243,25 +243,25 @@ func TestRTSPSourceMissingH264Params(t *testing.T) { pkts, err := enc.Encode([][]byte{{5}}, 0) // IDR require.NoError(t, err) - stream.WritePacketRTP(0, pkts[0]) + stream.WritePacketRTP(0, pkts[0], true) pkts, err = enc.Encode([][]byte{{7, 1, 2, 3}}, 0) // SPS require.NoError(t, err) - stream.WritePacketRTP(0, pkts[0]) + stream.WritePacketRTP(0, pkts[0], true) pkts, err = enc.Encode([][]byte{{8}}, 0) // PPS require.NoError(t, err) - stream.WritePacketRTP(0, pkts[0]) + stream.WritePacketRTP(0, pkts[0], true) pkts, err = enc.Encode([][]byte{{5, 1}}, 0) // IDR require.NoError(t, err) - stream.WritePacketRTP(0, pkts[0]) + stream.WritePacketRTP(0, pkts[0], true) time.Sleep(500 * time.Millisecond) pkts, err = enc.Encode([][]byte{{5, 2}}, 0) // IDR require.NoError(t, err) - stream.WritePacketRTP(0, pkts[0]) + stream.WritePacketRTP(0, pkts[0], true) }() return &base.Response{ @@ -286,17 +286,14 @@ func TestRTSPSourceMissingH264Params(t *testing.T) { defer p.close() received := make(chan struct{}) - decoder := &rtph264.Decoder{} - decoder.Init() c := gortsplib.Client{ - OnPacketRTP: func(trackID int, pkt *rtp.Packet) { - nalus, _, err := decoder.Decode(pkt) - if err != nil { + OnPacketRTP: func(ctx *gortsplib.ClientOnPacketRTPCtx) { + if ctx.H264NALUs == nil { return } - require.Equal(t, [][]byte{{0x05, 0x02}}, nalus) + require.Equal(t, [][]byte{{0x05, 0x02}}, ctx.H264NALUs) close(received) }, } diff --git a/internal/core/stream.go b/internal/core/stream.go index 036ee072..f5ea6066 100644 --- a/internal/core/stream.go +++ b/internal/core/stream.go @@ -4,7 +4,6 @@ import ( "sync" "github.com/aler9/gortsplib" - "github.com/pion/rtp" ) type streamNonRTSPReadersMap struct { @@ -36,12 +35,12 @@ func (m *streamNonRTSPReadersMap) remove(r reader) { delete(m.ma, r) } -func (m *streamNonRTSPReadersMap) forwardPacketRTP(trackID int, pkt *rtp.Packet) { +func (m *streamNonRTSPReadersMap) forwardPacketRTP(trackID int, data *data) { m.mutex.RLock() defer m.mutex.RUnlock() for c := range m.ma { - c.onReaderPacketRTP(trackID, pkt) + c.onReaderData(trackID, data) } } @@ -79,10 +78,10 @@ func (s *stream) readerRemove(r reader) { } } -func (s *stream) writePacketRTP(trackID int, pkt *rtp.Packet) { +func (s *stream) writeData(trackID int, data *data) { // forward to RTSP readers - s.rtspStream.WritePacketRTP(trackID, pkt) + s.rtspStream.WritePacketRTP(trackID, data.rtp, data.ptsEqualsDTS) // forward to non-RTSP readers - s.nonRTSPReaders.forwardPacketRTP(trackID, pkt) + s.nonRTSPReaders.forwardPacketRTP(trackID, data) } diff --git a/internal/hls/muxer_ts_generator.go b/internal/hls/muxer_ts_generator.go index 3a18f761..a67af983 100644 --- a/internal/hls/muxer_ts_generator.go +++ b/internal/hls/muxer_ts_generator.go @@ -14,16 +14,6 @@ const ( segmentMinAUCount = 100 ) -func idrPresent(nalus [][]byte) bool { - for _, nalu := range nalus { - typ := h264.NALUType(nalu[0] & 0x1F) - if typ == h264.NALUTypeIDR { - return true - } - } - return false -} - type writerFunc func(p []byte) (int, error) func (f writerFunc) Write(p []byte) (int, error) { @@ -93,7 +83,7 @@ func newMuxerTSGenerator( func (m *muxerTSGenerator) writeH264(pts time.Duration, nalus [][]byte) error { now := time.Now() - idrPresent := idrPresent(nalus) + idrPresent := h264.IDRPresent(nalus) if m.currentSegment == nil { // skip groups silently until we find one with a IDR