diff --git a/go.mod b/go.mod index 1037bedb..47e97c14 100644 --- a/go.mod +++ b/go.mod @@ -5,7 +5,7 @@ go 1.16 require ( github.com/alecthomas/template v0.0.0-20190718012654-fb15b899a751 // indirect github.com/alecthomas/units v0.0.0-20190924025748-f65c72e2690d // indirect - github.com/aler9/gortsplib v0.0.0-20211106122816-6e38851a096e + github.com/aler9/gortsplib v0.0.0-20211112212218-d205c0087835 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 2dbe8cfb..cc99c192 100644 --- a/go.sum +++ b/go.sum @@ -2,8 +2,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-20211106122816-6e38851a096e h1:qSjVAaIvJukmEuLxV0agmQ5KmBabBK+jzb+eNqG3Z+w= -github.com/aler9/gortsplib v0.0.0-20211106122816-6e38851a096e/go.mod h1:fyQrQyHo8QvdR/h357tkv1g36VesZlzEPsdAu2VrHHc= +github.com/aler9/gortsplib v0.0.0-20211112212218-d205c0087835 h1:GMW0OsdaXYUO67xhgtJUWll6gYQKAWiSDqcwhxHDCX8= +github.com/aler9/gortsplib v0.0.0-20211112212218-d205c0087835/go.mod h1:fyQrQyHo8QvdR/h357tkv1g36VesZlzEPsdAu2VrHHc= 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/api_test.go b/internal/core/api_test.go index cef51565..50264356 100644 --- a/internal/core/api_test.go +++ b/internal/core/api_test.go @@ -189,7 +189,9 @@ func TestAPIPathsList(t *testing.T) { require.NoError(t, err) func() { - source, err := gortsplib.DialPublish("rtsp://localhost:8554/mypath", + source := gortsplib.Client{} + + err = source.StartPublishing("rtsp://localhost:8554/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() @@ -200,7 +202,9 @@ func TestAPIPathsList(t *testing.T) { }() func() { - source, err := gortsplib.DialPublish("rtsps://localhost:8555/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsps://localhost:8555/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() @@ -249,13 +253,17 @@ func TestAPIList(t *testing.T) { switch ca { case "rtsp": - source, err := gortsplib.DialPublish("rtsp://localhost:8554/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsp://localhost:8554/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() case "rtsps": - source, err := gortsplib.DialPublish("rtsps://localhost:8555/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsps://localhost:8555/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() @@ -273,7 +281,9 @@ func TestAPIList(t *testing.T) { defer cnt1.close() case "hls": - source, err := gortsplib.DialPublish("rtsp://localhost:8554/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsp://localhost:8554/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() @@ -372,13 +382,17 @@ func TestAPIKick(t *testing.T) { switch ca { case "rtsp": - source, err := gortsplib.DialPublish("rtsp://localhost:8554/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsp://localhost:8554/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() case "rtsps": - source, err := gortsplib.DialPublish("rtsps://localhost:8555/mypath", + source := gortsplib.Client{} + + err := source.StartPublishing("rtsps://localhost:8555/mypath", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() diff --git a/internal/core/core_test.go b/internal/core/core_test.go index bc608c27..7ed6238f 100644 --- a/internal/core/core_test.go +++ b/internal/core/core_test.go @@ -12,7 +12,6 @@ import ( "github.com/aler9/gortsplib" "github.com/aler9/gortsplib/pkg/base" - "github.com/aler9/gortsplib/pkg/headers" psdp "github.com/pion/sdp/v3" "github.com/stretchr/testify/require" ) @@ -161,15 +160,17 @@ func TestCorePathAutoDeletion(t *testing.T) { defer p.close() func() { - conn, err := gortsplib.Dial("rtsp", "localhost:8554") + c := gortsplib.Client{} + + err := c.Start("rtsp", "localhost:8554") require.NoError(t, err) - defer conn.Close() + defer c.Close() if ca == "describe" { ur, err := base.ParseURL("rtsp://localhost:8554/mypath") require.NoError(t, err) - _, _, _, err = conn.Describe(ur) + _, _, _, err = c.Describe(ur) require.EqualError(t, err, "bad status code: 404 (Not Found)") } else { baseURL, err := base.ParseURL("rtsp://localhost:8554/mypath/") @@ -184,7 +185,7 @@ func TestCorePathAutoDeletion(t *testing.T) { Value: "trackID=0", }) - _, err = conn.Setup(headers.TransportModePlay, baseURL, track, 0, 0) + _, err = c.Setup(true, baseURL, track, 0, 0) require.EqualError(t, err, "bad status code: 404 (Not Found)") } }() @@ -219,7 +220,9 @@ func main() { panic(err) } - source, err := gortsplib.DialPublish( + source := gortsplib.Client{} + + err = source.StartPublishing( "rtsp://localhost:" + os.Getenv("RTSP_PORT") + "/" + os.Getenv("RTSP_PATH"), gortsplib.Tracks{track}) if err != nil { @@ -263,15 +266,17 @@ func main() { defer p1.close() func() { - conn, err := gortsplib.Dial("rtsp", "localhost:8554") + c := gortsplib.Client{} + + err := c.Start("rtsp", "localhost:8554") require.NoError(t, err) - defer conn.Close() + defer c.Close() if ca == "describe" || ca == "describe and setup" { ur, err := base.ParseURL("rtsp://localhost:8554/ondemand") require.NoError(t, err) - _, _, _, err = conn.Describe(ur) + _, _, _, err = c.Describe(ur) require.NoError(t, err) } @@ -288,7 +293,7 @@ func main() { Value: "trackID=0", }) - _, err = conn.Setup(headers.TransportModePlay, baseURL, track, 0, 0) + _, err = c.Setup(true, baseURL, track, 0, 0) require.NoError(t, err) } }() @@ -324,7 +329,9 @@ func TestCoreHotReloading(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - _, err = gortsplib.DialPublish( + c := gortsplib.Client{} + + err = c.StartPublishing( "rtsp://localhost:8554/test1", gortsplib.Tracks{track}) require.EqualError(t, err, "bad status code: 401 (Unauthorized)") @@ -342,7 +349,9 @@ func TestCoreHotReloading(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - conn, err := gortsplib.DialPublish( + conn := gortsplib.Client{} + + err = conn.StartPublishing( "rtsp://localhost:8554/test1", gortsplib.Tracks{track}) require.NoError(t, err) diff --git a/internal/core/hls_muxer.go b/internal/core/hls_muxer.go index 5604d95e..23e47469 100644 --- a/internal/core/hls_muxer.go +++ b/internal/core/hls_muxer.go @@ -497,11 +497,13 @@ func (m *hlsMuxer) onReaderAccepted() { m.log(logger.Info, "is converting into HLS") } -// onReaderFrame implements reader. -func (m *hlsMuxer) onReaderFrame(trackID int, streamType gortsplib.StreamType, payload []byte) { - if streamType == gortsplib.StreamTypeRTP { - m.ringBuffer.Push(hlsMuxerTrackIDPayloadPair{trackID, payload}) - } +// onReaderPacketRTP implements reader. +func (m *hlsMuxer) onReaderPacketRTP(trackID int, payload []byte) { + m.ringBuffer.Push(hlsMuxerTrackIDPayloadPair{trackID, payload}) +} + +// onReaderPacketRTCP implements reader. +func (m *hlsMuxer) onReaderPacketRTCP(trackID int, payload []byte) { } // onReaderAPIDescribe implements reader. diff --git a/internal/core/hls_source.go b/internal/core/hls_source.go index b5cb98a3..b4cb4661 100644 --- a/internal/core/hls_source.go +++ b/internal/core/hls_source.go @@ -123,12 +123,12 @@ func (s *hlsSource) runInner() bool { s.Log(logger.Info, "ready") stream = res.Stream - rtcpSenders = rtcpsenderset.New(tracks, stream.onFrame) + rtcpSenders = rtcpsenderset.New(tracks, stream.onPacketRTCP) return nil } - onFrame := func(isVideo bool, payload []byte) { + onPacket := func(isVideo bool, payload []byte) { var trackID int if isVideo { trackID = videoTrackID @@ -137,8 +137,8 @@ func (s *hlsSource) runInner() bool { } if stream != nil { - rtcpSenders.OnFrame(trackID, gortsplib.StreamTypeRTP, payload) - stream.onFrame(trackID, gortsplib.StreamTypeRTP, payload) + rtcpSenders.OnPacketRTP(trackID, payload) + stream.onPacketRTP(trackID, payload) } } @@ -146,7 +146,7 @@ func (s *hlsSource) runInner() bool { s.ur, s.fingerprint, onTracks, - onFrame, + onPacket, s, ) if err != nil { diff --git a/internal/core/hls_source_test.go b/internal/core/hls_source_test.go index 27b1d34c..d2c0b41d 100644 --- a/internal/core/hls_source_test.go +++ b/internal/core/hls_source_test.go @@ -6,7 +6,6 @@ import ( "io" "net" "net/http" - "sync/atomic" "testing" "time" @@ -132,28 +131,21 @@ func TestHLSSource(t *testing.T) { time.Sleep(1 * time.Second) - dest, err := gortsplib.DialRead("rtsp://localhost:8554/proxied") - require.NoError(t, err) - - rtcpRecv := int64(0) - readDone := make(chan struct{}) frameRecv := make(chan struct{}) - go func() { - defer close(readDone) - dest.ReadFrames(func(trackID int, streamType gortsplib.StreamType, payload []byte) { - if atomic.SwapInt64(&rtcpRecv, 1) == 0 { - } else { - require.Equal(t, gortsplib.StreamTypeRTP, streamType) - var pkt rtp.Packet - err := pkt.Unmarshal(payload) - require.NoError(t, err) - require.Equal(t, []byte{0x05}, pkt.Payload) - close(frameRecv) - } - }) - }() + + c := gortsplib.Client{ + OnPacketRTP: func(trackID int, payload []byte) { + var pkt rtp.Packet + err := pkt.Unmarshal(payload) + require.NoError(t, err) + require.Equal(t, []byte{0x05}, pkt.Payload) + close(frameRecv) + }, + } + + err = c.StartReading("rtsp://localhost:8554/proxied") + require.NoError(t, err) + defer c.Close() <-frameRecv - dest.Close() - <-readDone } diff --git a/internal/core/metrics_test.go b/internal/core/metrics_test.go index 4650d8f2..f41ff353 100644 --- a/internal/core/metrics_test.go +++ b/internal/core/metrics_test.go @@ -33,7 +33,9 @@ func TestMetrics(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - source, err := gortsplib.DialPublish("rtsp://localhost:8554/rtsp_path", + source := gortsplib.Client{} + + err = source.StartPublishing("rtsp://localhost:8554/rtsp_path", gortsplib.Tracks{track}) require.NoError(t, err) defer source.Close() diff --git a/internal/core/reader.go b/internal/core/reader.go index 526a781d..826205e7 100644 --- a/internal/core/reader.go +++ b/internal/core/reader.go @@ -1,13 +1,10 @@ package core -import ( - "github.com/aler9/gortsplib" -) - // reader is an entity that can read a stream. type reader interface { close() onReaderAccepted() - onReaderFrame(int, gortsplib.StreamType, []byte) + onReaderPacketRTP(int, []byte) + onReaderPacketRTCP(int, []byte) onReaderAPIDescribe() interface{} } diff --git a/internal/core/rtmp_conn.go b/internal/core/rtmp_conn.go index 78d3bb20..ef74a006 100644 --- a/internal/core/rtmp_conn.go +++ b/internal/core/rtmp_conn.go @@ -485,12 +485,12 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { return rres.Err } - rtcpSenders := rtcpsenderset.New(tracks, rres.Stream.onFrame) + rtcpSenders := rtcpsenderset.New(tracks, rres.Stream.onPacketRTCP) defer rtcpSenders.Close() - onFrame := func(trackID int, payload []byte) { - rtcpSenders.OnFrame(trackID, gortsplib.StreamTypeRTP, payload) - rres.Stream.onFrame(trackID, gortsplib.StreamTypeRTP, payload) + onPacketRTP := func(trackID int, payload []byte) { + rtcpSenders.OnPacketRTP(trackID, payload) + rres.Stream.onPacketRTP(trackID, payload) } for { @@ -503,7 +503,7 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { switch pkt.Type { case av.H264: if videoTrack == nil { - return fmt.Errorf("received an H264 frame, but track is not set up") + return fmt.Errorf("received an H264 packet, but track is not set up") } nalus, err := h264.DecodeAVCC(pkt.Data) @@ -543,12 +543,12 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { } for _, byts := range bytss { - onFrame(videoTrackID, byts) + onPacketRTP(videoTrackID, byts) } case av.AAC: if audioTrack == nil { - return fmt.Errorf("received an AAC frame, but track is not set up") + return fmt.Errorf("received an AAC packet, but track is not set up") } pkts, err := aacEncoder.Encode([][]byte{pkt.Data}, pkt.Time+pkt.CTime) @@ -566,7 +566,7 @@ func (c *rtmpConn) runPublish(ctx context.Context) error { } for _, byts := range bytss { - onFrame(audioTrackID, byts) + onPacketRTP(audioTrackID, byts) } } } @@ -592,11 +592,13 @@ func (c *rtmpConn) onReaderAccepted() { c.log(logger.Info, "is reading from path '%s'", c.path.Name()) } -// onReaderFrame implements reader. -func (c *rtmpConn) onReaderFrame(trackID int, streamType gortsplib.StreamType, payload []byte) { - if streamType == gortsplib.StreamTypeRTP { - c.ringBuffer.Push(rtmpConnTrackIDPayloadPair{trackID, payload}) - } +// onReaderPacketRTP implements reader. +func (c *rtmpConn) onReaderPacketRTP(trackID int, payload []byte) { + c.ringBuffer.Push(rtmpConnTrackIDPayloadPair{trackID, payload}) +} + +// onReaderPacketRTCP implements reader. +func (c *rtmpConn) onReaderPacketRTCP(trackID int, payload []byte) { } // onReaderAPIDescribe implements reader. diff --git a/internal/core/rtmp_source.go b/internal/core/rtmp_source.go index d34f0201..c2698771 100644 --- a/internal/core/rtmp_source.go +++ b/internal/core/rtmp_source.go @@ -163,12 +163,12 @@ func (s *rtmpSource) runInner() bool { s.parent.OnSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{Source: s}) }() - rtcpSenders := rtcpsenderset.New(tracks, res.Stream.onFrame) + rtcpSenders := rtcpsenderset.New(tracks, res.Stream.onPacketRTCP) defer rtcpSenders.Close() - onFrame := func(trackID int, payload []byte) { - rtcpSenders.OnFrame(trackID, gortsplib.StreamTypeRTP, payload) - res.Stream.onFrame(trackID, gortsplib.StreamTypeRTP, payload) + onPacketRTP := func(trackID int, payload []byte) { + rtcpSenders.OnPacketRTP(trackID, payload) + res.Stream.onPacketRTP(trackID, payload) } for { @@ -181,7 +181,7 @@ func (s *rtmpSource) runInner() bool { switch pkt.Type { case av.H264: if videoTrack == nil { - return fmt.Errorf("received an H264 frame, but track is not set up") + return fmt.Errorf("received an H264 packet, but track is not set up") } nalus, err := h264.DecodeAVCC(pkt.Data) @@ -216,12 +216,12 @@ func (s *rtmpSource) runInner() bool { } for _, byts := range bytss { - onFrame(videoTrackID, byts) + onPacketRTP(videoTrackID, byts) } case av.AAC: if audioTrack == nil { - return fmt.Errorf("received an AAC frame, but track is not set up") + return fmt.Errorf("received an AAC packet, but track is not set up") } pkts, err := aacEncoder.Encode([][]byte{pkt.Data}, pkt.Time+pkt.CTime) @@ -239,7 +239,7 @@ func (s *rtmpSource) runInner() bool { } for _, byts := range bytss { - onFrame(audioTrackID, byts) + onPacketRTP(audioTrackID, byts) } } } diff --git a/internal/core/rtsp_server.go b/internal/core/rtsp_server.go index f4c11af5..5e6adf0b 100644 --- a/internal/core/rtsp_server.go +++ b/internal/core/rtsp_server.go @@ -117,6 +117,7 @@ func newRTSPServer( WriteTimeout: time.Duration(writeTimeout), ReadBufferCount: readBufferCount, ReadBufferSize: readBufferSize, + RTSPAddress: address, } if useUDP { @@ -139,7 +140,7 @@ func newRTSPServer( s.srv.TLSConfig = &tls.Config{Certificates: []tls.Certificate{cert}} } - err := s.srv.Start(address) + err := s.srv.Start() if err != nil { return nil, err } @@ -375,12 +376,20 @@ func (s *rtspServer) OnPause(ctx *gortsplib.ServerHandlerOnPauseCtx) (*base.Resp return se.onPause(ctx) } -// OnFrame implements gortsplib.ServerHandlerOnFrame. -func (s *rtspServer) OnFrame(ctx *gortsplib.ServerHandlerOnFrameCtx) { +// OnPacketRTP implements gortsplib.ServerHandlerOnPacket. +func (s *rtspServer) OnPacketRTP(ctx *gortsplib.ServerHandlerOnPacketRTPCtx) { s.mutex.RLock() se := s.sessions[ctx.Session] s.mutex.RUnlock() - se.onFrame(ctx) + se.onPacketRTP(ctx) +} + +// OnPacketRTCP implements gortsplib.ServerHandlerOnPacket. +func (s *rtspServer) OnPacketRTCP(ctx *gortsplib.ServerHandlerOnPacketRTCPCtx) { + s.mutex.RLock() + se := s.sessions[ctx.Session] + s.mutex.RUnlock() + se.onPacketRTCP(ctx) } // onAPISessionsList is called by api and metrics. diff --git a/internal/core/rtsp_server_test.go b/internal/core/rtsp_server_test.go index 5d64bebb..66c280b8 100644 --- a/internal/core/rtsp_server_test.go +++ b/internal/core/rtsp_server_test.go @@ -6,7 +6,6 @@ import ( "time" "github.com/aler9/gortsplib" - "github.com/aler9/gortsplib/pkg/base" "github.com/stretchr/testify/require" ) @@ -204,7 +203,9 @@ func TestRTSPServerAuth(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - source, err := gortsplib.DialPublish( + source := gortsplib.Client{} + + err = source.StartPublishing( "rtsp://testuser:test%21%24%28%29%2A%2B.%3B%3C%3D%3E%5B%5D%5E_-%7B%7D@127.0.0.1:8554/test/stream", gortsplib.Tracks{track}) require.NoError(t, err) @@ -276,7 +277,9 @@ func TestRTSPServerAuth(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - source, err := gortsplib.DialPublish( + source := gortsplib.Client{} + + err = source.StartPublishing( "rtsp://testuser:testpass@127.0.0.1:8554/test/stream", gortsplib.Tracks{track}) require.NoError(t, err) @@ -320,7 +323,9 @@ func TestRTSPServerAuthFail(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - _, err = gortsplib.DialPublish( + c := gortsplib.Client{} + + err = c.StartPublishing( "rtsp://"+ca.user+":"+ca.pass+"@localhost:8554/test/stream", gortsplib.Tracks{track}, ) @@ -359,7 +364,9 @@ func TestRTSPServerAuthFail(t *testing.T) { require.Equal(t, true, ok) defer p.close() - _, err := gortsplib.DialRead( + c := gortsplib.Client{} + + err := c.StartReading( "rtsp://" + ca.user + ":" + ca.pass + "@localhost:8554/test/stream", ) require.EqualError(t, err, "bad status code: 401 (Unauthorized)") @@ -379,7 +386,9 @@ func TestRTSPServerAuthFail(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - _, err = gortsplib.DialPublish( + c := gortsplib.Client{} + + err = c.StartPublishing( "rtsp://localhost:8554/test/stream", gortsplib.Tracks{track}, ) @@ -410,12 +419,16 @@ func TestRTSPServerPublisherOverride(t *testing.T) { &gortsplib.TrackConfigH264{SPS: []byte{0x01, 0x02, 0x03, 0x04}, PPS: []byte{0x01, 0x02, 0x03, 0x04}}) require.NoError(t, err) - s1, err := gortsplib.DialPublish("rtsp://localhost:8554/teststream", + s1 := gortsplib.Client{} + + err = s1.StartPublishing("rtsp://localhost:8554/teststream", gortsplib.Tracks{track}) require.NoError(t, err) defer s1.Close() - s2, err := gortsplib.DialPublish("rtsp://localhost:8554/teststream", + s2 := gortsplib.Client{} + + err = s2.StartPublishing("rtsp://localhost:8554/teststream", gortsplib.Tracks{track}) if ca == "enabled" { require.NoError(t, err) @@ -424,27 +437,24 @@ func TestRTSPServerPublisherOverride(t *testing.T) { require.Error(t, err) } - d1, err := gortsplib.DialRead("rtsp://localhost:8554/teststream") - require.NoError(t, err) - defer d1.Close() - - readDone := make(chan struct{}) frameRecv := make(chan struct{}) - go func() { - defer close(readDone) - d1.ReadFrames(func(trackID int, streamType base.StreamType, payload []byte) { - if streamType == gortsplib.StreamTypeRTP { - if ca == "enabled" { - require.Equal(t, []byte{0x05, 0x06, 0x07, 0x08}, payload) - } else { - require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, payload) - } - close(frameRecv) + + c := gortsplib.Client{ + OnPacketRTP: func(trackID int, payload []byte) { + if ca == "enabled" { + require.Equal(t, []byte{0x05, 0x06, 0x07, 0x08}, payload) + } else { + require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, payload) } - }) - }() + close(frameRecv) + }, + } + + err = c.StartReading("rtsp://localhost:8554/teststream") + require.NoError(t, err) + defer c.Close() - err = s1.WriteFrame(0, gortsplib.StreamTypeRTP, + err = s1.WritePacketRTP(0, []byte{0x01, 0x02, 0x03, 0x04}) if ca == "enabled" { require.Error(t, err) @@ -453,15 +463,12 @@ func TestRTSPServerPublisherOverride(t *testing.T) { } if ca == "enabled" { - err = s2.WriteFrame(0, gortsplib.StreamTypeRTP, + err = s2.WritePacketRTP(0, []byte{0x05, 0x06, 0x07, 0x08}) require.NoError(t, err) } <-frameRecv - - d1.Close() - <-readDone }) } } diff --git a/internal/core/rtsp_session.go b/internal/core/rtsp_session.go index 3e696112..e30ef983 100644 --- a/internal/core/rtsp_session.go +++ b/internal/core/rtsp_session.go @@ -335,9 +335,14 @@ func (s *rtspSession) onReaderAccepted() { s.ss.SetuppedTransport()) } -// onReaderFrame implements reader. -func (s *rtspSession) onReaderFrame(trackID int, streamType gortsplib.StreamType, payload []byte) { - s.ss.WriteFrame(trackID, streamType, payload) +// onReaderPacketRTP implements reader. +func (s *rtspSession) onReaderPacketRTP(trackID int, payload []byte) { + s.ss.WritePacketRTP(trackID, payload) +} + +// onReaderPacketRTCP implements reader. +func (s *rtspSession) onReaderPacketRTCP(trackID int, payload []byte) { + s.ss.WritePacketRTCP(trackID, payload) } // onReaderAPIDescribe implements reader. @@ -384,11 +389,20 @@ func (s *rtspSession) onPublisherAccepted(tracksLen int) { s.ss.SetuppedTransport()) } -// onFrame is called by rtspServer. -func (s *rtspSession) onFrame(ctx *gortsplib.ServerHandlerOnFrameCtx) { +// onPacketRTP is called by rtspServer. +func (s *rtspSession) onPacketRTP(ctx *gortsplib.ServerHandlerOnPacketRTPCtx) { + if s.ss.State() != gortsplib.ServerSessionStatePublish { + return + } + + s.stream.onPacketRTP(ctx.TrackID, ctx.Payload) +} + +// onPacketRTCP is called by rtspServer. +func (s *rtspSession) onPacketRTCP(ctx *gortsplib.ServerHandlerOnPacketRTCPCtx) { if s.ss.State() != gortsplib.ServerSessionStatePublish { return } - s.stream.onFrame(ctx.TrackID, ctx.StreamType, ctx.Payload) + s.stream.onPacketRTCP(ctx.TrackID, ctx.Payload) } diff --git a/internal/core/rtsp_source.go b/internal/core/rtsp_source.go index 46061e0f..b605d6ee 100644 --- a/internal/core/rtsp_source.go +++ b/internal/core/rtsp_source.go @@ -118,7 +118,6 @@ func (s *rtspSource) runInner() bool { s.log(logger.Debug, "connecting") tlsConfig := &tls.Config{} - if s.fingerprint != "" { tlsConfig.InsecureSkipVerify = true tlsConfig.VerifyConnection = func(cs tls.ConnectionState) error { @@ -136,7 +135,7 @@ func (s *rtspSource) runInner() bool { } } - client := &gortsplib.Client{ + c := &gortsplib.Client{ Transport: s.proto.Transport, TLSConfig: tlsConfig, ReadTimeout: time.Duration(s.readTimeout), @@ -152,63 +151,78 @@ func (s *rtspSource) runInner() bool { }, } - innerCtx, innerCtxCancel := context.WithCancel(context.Background()) - - var conn *gortsplib.ClientConn - var err error - dialDone := make(chan struct{}) - go func() { - defer close(dialDone) - conn, err = client.DialReadContext(innerCtx, s.ur) - }() - - select { - case <-s.ctx.Done(): - innerCtxCancel() - <-dialDone - return false - - case <-dialDone: - innerCtxCancel() - } - + u, err := base.ParseURL(s.ur) if err != nil { s.log(logger.Info, "ERR: %s", err) return true } - res := s.parent.onSourceStaticSetReady(pathSourceStaticSetReadyReq{ - Source: s, - Tracks: conn.Tracks(), - }) - if res.Err != nil { - s.log(logger.Info, "ERR: %s", res.Err) + err = c.Start(u.Scheme, u.Host) + if err != nil { + s.log(logger.Info, "ERR: %s", err) return true } - s.log(logger.Info, "ready") - - defer func() { - s.parent.OnSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{Source: s}) - }() - readErr := make(chan error) go func() { - readErr <- conn.ReadFrames(func(trackID int, streamType gortsplib.StreamType, payload []byte) { - res.Stream.onFrame(trackID, streamType, payload) - }) + readErr <- func() error { + _, err = c.Options(u) + if err != nil { + return err + } + + tracks, baseURL, _, err := c.Describe(u) + if err != nil { + return err + } + + for _, t := range tracks { + _, err := c.Setup(true, baseURL, t, 0, 0) + if err != nil { + panic(err) + } + } + + res := s.parent.onSourceStaticSetReady(pathSourceStaticSetReadyReq{ + Source: s, + Tracks: c.Tracks(), + }) + if res.Err != nil { + return res.Err + } + + s.log(logger.Info, "ready") + + defer func() { + s.parent.OnSourceStaticSetNotReady(pathSourceStaticSetNotReadyReq{Source: s}) + }() + + c.OnPacketRTP = func(trackID int, payload []byte) { + res.Stream.onPacketRTP(trackID, payload) + } + + c.OnPacketRTCP = func(trackID int, payload []byte) { + res.Stream.onPacketRTCP(trackID, payload) + } + + _, err = c.Play(nil) + if err != nil { + return err + } + + return c.Wait() + }() }() select { - case <-s.ctx.Done(): - conn.Close() - <-readErr - return false - case err := <-readErr: s.log(logger.Info, "ERR: %s", err) - conn.Close() return true + + case <-s.ctx.Done(): + c.Close() + <-readErr + return false } } diff --git a/internal/core/rtsp_source_test.go b/internal/core/rtsp_source_test.go index 5e0b8580..b1fa7020 100644 --- a/internal/core/rtsp_source_test.go +++ b/internal/core/rtsp_source_test.go @@ -59,7 +59,7 @@ func (sh *testServer) OnSetup(ctx *gortsplib.ServerHandlerOnSetupCtx) (*base.Res func (sh *testServer) OnPlay(ctx *gortsplib.ServerHandlerOnPlayCtx) (*base.Response, error) { go func() { time.Sleep(1 * time.Second) - sh.stream.WriteFrame(0, gortsplib.StreamTypeRTP, []byte{0x01, 0x02, 0x03, 0x04}) + sh.stream.WritePacketRTP(0, []byte{0x01, 0x02, 0x03, 0x04}) }() return &base.Response{ @@ -75,7 +75,8 @@ func TestRTSPSource(t *testing.T) { } { t.Run(source, func(t *testing.T) { s := gortsplib.Server{ - Handler: &testServer{user: "testuser", pass: "testpass"}, + Handler: &testServer{user: "testuser", pass: "testpass"}, + RTSPAddress: "127.0.0.1:8555", } switch source { @@ -98,7 +99,7 @@ func TestRTSPSource(t *testing.T) { s.TLSConfig = &tls.Config{Certificates: []tls.Certificate{cert}} } - err := s.Start("127.0.0.1:8555") + err := s.Start() require.NoError(t, err) defer s.Wait() defer s.Close() @@ -123,32 +124,31 @@ func TestRTSPSource(t *testing.T) { time.Sleep(1 * time.Second) - conn, err := gortsplib.DialRead("rtsp://127.0.0.1:8554/proxied") - require.NoError(t, err) - - readDone := make(chan struct{}) received := make(chan struct{}) - go func() { - defer close(readDone) - conn.ReadFrames(func(trackID int, streamType gortsplib.StreamType, payload []byte) { - if streamType == gortsplib.StreamTypeRTP { - require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, payload) - close(received) - } - }) - }() + + c := gortsplib.Client{ + OnPacketRTP: func(trackID int, payload []byte) { + require.Equal(t, []byte{0x01, 0x02, 0x03, 0x04}, payload) + close(received) + }, + } + + err = c.StartReading("rtsp://127.0.0.1:8554/proxied") + require.NoError(t, err) + defer c.Close() <-received - conn.Close() - <-readDone }) } } func TestRTSPSourceNoPassword(t *testing.T) { done := make(chan struct{}) - s := gortsplib.Server{Handler: &testServer{user: "testuser", done: done}} - err := s.Start("127.0.0.1:8555") + s := gortsplib.Server{ + Handler: &testServer{user: "testuser", done: done}, + RTSPAddress: "127.0.0.1:8555", + } + err := s.Start() require.NoError(t, err) defer s.Wait() defer s.Close() diff --git a/internal/core/stream.go b/internal/core/stream.go index cbba9fdc..fd517375 100644 --- a/internal/core/stream.go +++ b/internal/core/stream.go @@ -35,12 +35,21 @@ func (m *streamNonRTSPReadersMap) remove(r reader) { delete(m.ma, r) } -func (m *streamNonRTSPReadersMap) forwardFrame(trackID int, streamType gortsplib.StreamType, payload []byte) { +func (m *streamNonRTSPReadersMap) forwardPacketRTP(trackID int, payload []byte) { m.mutex.RLock() defer m.mutex.RUnlock() for c := range m.ma { - c.onReaderFrame(trackID, streamType, payload) + c.onReaderPacketRTP(trackID, payload) + } +} + +func (m *streamNonRTSPReadersMap) forwardPacketRTCP(trackID int, payload []byte) { + m.mutex.RLock() + defer m.mutex.RUnlock() + + for c := range m.ma { + c.onReaderPacketRTCP(trackID, payload) } } @@ -78,10 +87,18 @@ func (s *stream) readerRemove(r reader) { } } -func (s *stream) onFrame(trackID int, streamType gortsplib.StreamType, payload []byte) { +func (s *stream) onPacketRTP(trackID int, payload []byte) { + // forward to RTSP readers + s.rtspStream.WritePacketRTP(trackID, payload) + + // forward to non-RTSP readers + s.nonRTSPReaders.forwardPacketRTP(trackID, payload) +} + +func (s *stream) onPacketRTCP(trackID int, payload []byte) { // forward to RTSP readers - s.rtspStream.WriteFrame(trackID, streamType, payload) + s.rtspStream.WritePacketRTCP(trackID, payload) // forward to non-RTSP readers - s.nonRTSPReaders.forwardFrame(trackID, streamType, payload) + s.nonRTSPReaders.forwardPacketRTCP(trackID, payload) } diff --git a/internal/hls/client.go b/internal/hls/client.go index 2f6c45f2..fed06c5f 100644 --- a/internal/hls/client.go +++ b/internal/hls/client.go @@ -133,9 +133,9 @@ type clientVideoProcessorData struct { } type clientVideoProcessor struct { - ctx context.Context - onTrack func(*gortsplib.Track) error - onFrame func([]byte) + ctx context.Context + onTrack func(*gortsplib.Track) error + onPacket func([]byte) queue chan clientVideoProcessorData sps []byte @@ -147,13 +147,13 @@ type clientVideoProcessor struct { func newClientVideoProcessor( ctx context.Context, onTrack func(*gortsplib.Track) error, - onFrame func([]byte), + onPacket func([]byte), ) *clientVideoProcessor { p := &clientVideoProcessor{ - ctx: ctx, - onTrack: onTrack, - onFrame: onFrame, - queue: make(chan clientVideoProcessorData, clientQueueSize), + ctx: ctx, + onTrack: onTrack, + onPacket: onPacket, + queue: make(chan clientVideoProcessorData, clientQueueSize), } return p @@ -259,7 +259,7 @@ func (p *clientVideoProcessor) doProcess( } for _, byts := range bytss { - p.onFrame(byts) + p.onPacket(byts) } return nil @@ -289,9 +289,9 @@ type clientAudioProcessorData struct { } type clientAudioProcessor struct { - ctx context.Context - onTrack func(*gortsplib.Track) error - onFrame func([]byte) + ctx context.Context + onTrack func(*gortsplib.Track) error + onPacket func([]byte) queue chan clientAudioProcessorData conf *gortsplib.TrackConfigAAC @@ -302,13 +302,13 @@ type clientAudioProcessor struct { func newClientAudioProcessor( ctx context.Context, onTrack func(*gortsplib.Track) error, - onFrame func([]byte), + onPacket func([]byte), ) *clientAudioProcessor { p := &clientAudioProcessor{ - ctx: ctx, - onTrack: onTrack, - onFrame: onFrame, - queue: make(chan clientAudioProcessorData, clientQueueSize), + ctx: ctx, + onTrack: onTrack, + onPacket: onPacket, + queue: make(chan clientAudioProcessorData, clientQueueSize), } return p @@ -392,7 +392,7 @@ func (p *clientAudioProcessor) doProcess( } for _, byts := range bytss { - p.onFrame(byts) + p.onPacket(byts) } return nil @@ -426,7 +426,7 @@ type ClientParent interface { // Client is a HLS client. type Client struct { onTracks func(*gortsplib.Track, *gortsplib.Track) error - onFrame func(bool, []byte) + onPacket func(bool, []byte) parent ClientParent ctx context.Context @@ -462,7 +462,7 @@ func NewClient( primaryPlaylistURLStr string, fingerprint string, onTracks func(*gortsplib.Track, *gortsplib.Track) error, - onFrame func(bool, []byte), + onPacket func(bool, []byte), parent ClientParent, ) (*Client, error) { primaryPlaylistURL, err := url.Parse(primaryPlaylistURLStr) @@ -493,7 +493,7 @@ func NewClient( c := &Client{ onTracks: onTracks, - onFrame: onFrame, + onPacket: onPacket, parent: parent, ctx: ctx, ctxCancel: ctxCancel, @@ -546,7 +546,7 @@ func (c *Client) runInner() error { c.videoProc = newClientVideoProcessor( innerCtx, c.onVideoTrack, - c.onVideoFrame) + c.onVideoPacket) go func() { errChan <- c.videoProc.run() }() } @@ -555,7 +555,7 @@ func (c *Client) runInner() error { c.audioProc = newClientAudioProcessor( innerCtx, c.onAudioTrack, - c.onAudioFrame) + c.onAudioPacket) go func() { errChan <- c.audioProc.run() }() } @@ -924,16 +924,16 @@ func (c *Client) initializeTracks() error { return c.onTracks(c.videoTrack, c.audioTrack) } -func (c *Client) onVideoFrame(payload []byte) { +func (c *Client) onVideoPacket(payload []byte) { c.tracksMutex.RLock() defer c.tracksMutex.RUnlock() - c.onFrame(true, payload) + c.onPacket(true, payload) } -func (c *Client) onAudioFrame(payload []byte) { +func (c *Client) onAudioPacket(payload []byte) { c.tracksMutex.RLock() defer c.tracksMutex.RUnlock() - c.onFrame(false, payload) + c.onPacket(false, payload) } diff --git a/internal/hls/client_test.go b/internal/hls/client_test.go index 8bcf917a..9da94f88 100644 --- a/internal/hls/client_test.go +++ b/internal/hls/client_test.go @@ -190,7 +190,7 @@ func TestClient(t *testing.T) { require.NoError(t, err) defer ts.close() - frameRecv := make(chan struct{}) + packetRecv := make(chan struct{}) prefix := "http" if mode == "tls" { @@ -206,13 +206,13 @@ func TestClient(t *testing.T) { func(isVideo bool, byts []byte) { require.Equal(t, true, isVideo) require.Equal(t, byte(0x05), byts[12]) - close(frameRecv) + close(packetRecv) }, testClientParent{}, ) require.NoError(t, err) - <-frameRecv + <-packetRecv c.Close() c.Wait() diff --git a/internal/rtcpsenderset/rtcpsenderset.go b/internal/rtcpsenderset/rtcpsenderset.go index 58980b16..dc2fc2a5 100644 --- a/internal/rtcpsenderset/rtcpsenderset.go +++ b/internal/rtcpsenderset/rtcpsenderset.go @@ -9,8 +9,8 @@ import ( // RTCPSenderSet is a set of RTCP senders. type RTCPSenderSet struct { - onFrame func(int, gortsplib.StreamType, []byte) - senders []*rtcpsender.RTCPSender + onPacketRTCP func(int, []byte) + senders []*rtcpsender.RTCPSender // in terminate chan struct{} @@ -22,12 +22,12 @@ type RTCPSenderSet struct { // New allocates a RTCPSenderSet. func New( tracks gortsplib.Tracks, - onFrame func(int, gortsplib.StreamType, []byte), + onPacketRTCP func(int, []byte), ) *RTCPSenderSet { s := &RTCPSenderSet{ - onFrame: onFrame, - terminate: make(chan struct{}), - done: make(chan struct{}), + onPacketRTCP: onPacketRTCP, + terminate: make(chan struct{}), + done: make(chan struct{}), } s.senders = make([]*rtcpsender.RTCPSender, len(tracks)) @@ -61,7 +61,7 @@ func (s *RTCPSenderSet) run() { for i, sender := range s.senders { r := sender.Report(now) if r != nil { - s.onFrame(i, gortsplib.StreamTypeRTCP, r) + s.onPacketRTCP(i, r) } } @@ -71,7 +71,12 @@ func (s *RTCPSenderSet) run() { } } -// OnFrame sends a frame to the senders. -func (s *RTCPSenderSet) OnFrame(trackID int, streamType gortsplib.StreamType, f []byte) { - s.senders[trackID].ProcessFrame(time.Now(), streamType, f) +// OnPacketRTP sends a RTP packet to the senders. +func (s *RTCPSenderSet) OnPacketRTP(trackID int, payload []byte) { + s.senders[trackID].ProcessPacketRTP(time.Now(), payload) +} + +// OnPacketRTCP sends a RTCP packet to the senders. +func (s *RTCPSenderSet) OnPacketRTCP(trackID int, payload []byte) { + s.senders[trackID].ProcessPacketRTCP(time.Now(), payload) }