From 1b9dfbd367a572222ee0a924ba146ac5598fb984 Mon Sep 17 00:00:00 2001 From: Alessandro Ros Date: Thu, 15 May 2025 14:23:03 +0200 Subject: [PATCH] rtmp: support connecting to sources that require standard credentials (#4530) --- internal/core/api_test.go | 47 +- internal/core/metrics_test.go | 30 +- internal/core/path_test.go | 54 +- internal/protocols/rtmp/client.go | 492 +++++++++++++ internal/protocols/rtmp/client_test.go | 421 +++++++++++ internal/protocols/rtmp/conn.go | 617 +---------------- internal/protocols/rtmp/conn_test.go | 542 --------------- internal/protocols/rtmp/from_stream.go | 2 +- internal/protocols/rtmp/from_stream_test.go | 10 +- internal/protocols/rtmp/reader.go | 6 +- internal/protocols/rtmp/reader_test.go | 10 +- internal/protocols/rtmp/server_conn.go | 509 ++++++++++++++ internal/protocols/rtmp/server_conn_test.go | 728 ++++++++++++++++++++ internal/protocols/rtmp/writer.go | 6 +- internal/protocols/rtmp/writer_test.go | 10 +- internal/servers/rtmp/conn.go | 47 +- internal/servers/rtmp/server_test.go | 72 +- internal/staticsources/rtmp/source.go | 71 +- internal/staticsources/rtmp/source_test.go | 138 ++-- 19 files changed, 2408 insertions(+), 1404 deletions(-) create mode 100644 internal/protocols/rtmp/client.go create mode 100644 internal/protocols/rtmp/client_test.go delete mode 100644 internal/protocols/rtmp/conn_test.go create mode 100644 internal/protocols/rtmp/server_conn.go create mode 100644 internal/protocols/rtmp/server_conn_test.go diff --git a/internal/core/api_test.go b/internal/core/api_test.go index 6c2eaae6..7605d453 100644 --- a/internal/core/api_test.go +++ b/internal/core/api_test.go @@ -8,7 +8,6 @@ import ( "crypto/tls" "encoding/json" "io" - "net" "net/http" "net/url" "os" @@ -418,27 +417,28 @@ func TestAPIProtocolListGet(t *testing.T) { port = "1936" } - u, err := url.Parse("rtmp://127.0.0.1:" + port + "/mypath?key=val") - require.NoError(t, err) + var rawURL string - nconn, err := func() (net.Conn, error) { - if ca == "rtmp" { - return net.Dial("tcp", u.Host) - } - return tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - }() - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: true, + if ca == "rtmps" { + rawURL = "rtmps://" + } else { + rawURL = "rtmp://" } - err = conn.Initialize() + + rawURL += "127.0.0.1:" + port + "/mypath?key=val" + + u, err := url.Parse(rawURL) require.NoError(t, err) + conn := &rtmp.Client{ + URL: u, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + Publish: true, + } + err = conn.Initialize(context.Background()) + require.NoError(t, err) + defer conn.Close() + w := &rtmp.Writer{ Conn: conn, VideoTrack: test.FormatH264, @@ -1013,18 +1013,13 @@ func TestAPIProtocolKick(t *testing.T) { u, err := url.Parse("rtmp://localhost:1935/mypath") require.NoError(t, err) - nconn, err := net.Dial("tcp", u.Host) - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, + conn := &rtmp.Client{ URL: u, Publish: true, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() w := &rtmp.Writer{ Conn: conn, diff --git a/internal/core/metrics_test.go b/internal/core/metrics_test.go index 2a6191c6..48017dea 100644 --- a/internal/core/metrics_test.go +++ b/internal/core/metrics_test.go @@ -5,7 +5,6 @@ import ( "context" "crypto/tls" "io" - "net" "net/http" "net/url" "os" @@ -196,21 +195,17 @@ webrtc_sessions_bytes_sent 0 go func() { defer wg.Done() + u, err := url.Parse("rtmp://localhost:1935/rtmp_path") require.NoError(t, err) - nconn, err := net.Dial("tcp", u.Host) - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, + conn := &rtmp.Client{ URL: u, Publish: true, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() w := &rtmp.Writer{ Conn: conn, @@ -227,21 +222,18 @@ webrtc_sessions_bytes_sent 0 go func() { defer wg.Done() + u, err := url.Parse("rtmps://localhost:1936/rtmps_path") require.NoError(t, err) - nconn, err := tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - require.NoError(t, err) - defer nconn.Close() //nolint:errcheck - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: true, + conn := &rtmp.Client{ + URL: u, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + Publish: true, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() w := &rtmp.Writer{ Conn: conn, diff --git a/internal/core/path_test.go b/internal/core/path_test.go index 7d1cf5dd..36e21aa9 100644 --- a/internal/core/path_test.go +++ b/internal/core/path_test.go @@ -213,18 +213,13 @@ func TestPathRunOnConnect(t *testing.T) { u, err := url.Parse("rtmp://127.0.0.1:1935/test") require.NoError(t, err) - nconn, err := net.Dial("tcp", u.Host) - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, + conn := &rtmp.Client{ URL: u, Publish: true, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() case "rtmps": connType = "rtmpsConn" @@ -232,18 +227,14 @@ func TestPathRunOnConnect(t *testing.T) { u, err := url.Parse("rtmps://127.0.0.1:1936/test") require.NoError(t, err) - nconn, err := tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - require.NoError(t, err) - defer nconn.Close() //nolint:errcheck - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: true, + conn := &rtmp.Client{ + URL: u, + Publish: true, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() case "srt": connType = "srtConn" @@ -448,18 +439,13 @@ func TestPathRunOnRead(t *testing.T) { u, err := url.Parse("rtmp://127.0.0.1:1935/test?query=value") require.NoError(t, err) - nconn, err := net.Dial("tcp", u.Host) - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, + conn := &rtmp.Client{ URL: u, Publish: false, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() r := &rtmp.Reader{ Conn: conn, @@ -471,18 +457,14 @@ func TestPathRunOnRead(t *testing.T) { u, err := url.Parse("rtmps://127.0.0.1:1936/test?query=value") require.NoError(t, err) - nconn, err := tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - require.NoError(t, err) - defer nconn.Close() //nolint:errcheck - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: false, + conn := &rtmp.Client{ + URL: u, + Publish: false, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, } - err = conn.Initialize() + err = conn.Initialize(context.Background()) require.NoError(t, err) + defer conn.Close() go func() { for i := uint16(0); i < 3; i++ { diff --git a/internal/protocols/rtmp/client.go b/internal/protocols/rtmp/client.go new file mode 100644 index 00000000..b1c28f34 --- /dev/null +++ b/internal/protocols/rtmp/client.go @@ -0,0 +1,492 @@ +// Package rtmp provides RTMP connectivity. +package rtmp + +import ( + "context" + ctls "crypto/tls" + "errors" + "fmt" + "net" + "net/url" + "strings" + + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" + "github.com/google/uuid" +) + +var errAuth = errors.New("auth") + +func resultIsOK1(res *message.CommandAMF0) bool { + if len(res.Arguments) < 2 { + return false + } + + ma, ok := objectOrArray(res.Arguments[1]) + if !ok { + return false + } + + v, ok := ma.Get("level") + if !ok { + return false + } + + return (v == "status") +} + +func resultIsOK2(res *message.CommandAMF0) bool { + if len(res.Arguments) < 2 { + return false + } + + v, ok := res.Arguments[1].(float64) + if !ok { + return false + } + + return v == 1 +} + +func splitPath(u *url.URL) (string, string) { + nu := *u + nu.ForceQuery = false + pathsegs := strings.Split(nu.RequestURI(), "/") + + var app string + var streamKey string + + switch { + case len(pathsegs) == 2: + app = pathsegs[1] + + case len(pathsegs) == 3: + app = pathsegs[1] + streamKey = pathsegs[2] + + case len(pathsegs) > 3: + app = strings.Join(pathsegs[1:3], "/") + streamKey = strings.Join(pathsegs[3:], "/") + } + + return app, streamKey +} + +func getTcURL(u *url.URL) string { + app, _ := splitPath(u) + nu, _ := url.Parse(u.String()) // perform a deep copy + nu.RawQuery = "" + nu.Path = "/" + return nu.String() + app +} + +func readCommand(mrw *message.ReadWriter) (*message.CommandAMF0, error) { + for { + msg, err := mrw.Read() + if err != nil { + return nil, err + } + + if cmd, ok := msg.(*message.CommandAMF0); ok { + return cmd, nil + } + } +} + +func readCommandResult( + mrw *message.ReadWriter, + commandID int, +) (*message.CommandAMF0, error) { + for { + msg, err := mrw.Read() + if err != nil { + return nil, err + } + + if cmd, ok := msg.(*message.CommandAMF0); ok { + if cmd.CommandID == commandID || cmd.CommandID == 0 { + return cmd, nil + } + } + } +} + +type dialer interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) +} + +// Client is a client-side RTMP connection. +type Client struct { + URL *url.URL + TLSConfig *ctls.Config + Publish bool + + nconn net.Conn + bc *bytecounter.ReadWriter + mrw *message.ReadWriter + authState int + authSalt string + authChallenge string +} + +// Initialize initializes Client. +func (c *Client) Initialize(ctx context.Context) error { + for { + err := c.initialize2(ctx) + if errors.Is(err, errAuth) { + c.authState++ + continue + } + return err + } +} + +func (c *Client) initialize2(ctx context.Context) error { + var dial dialer + if c.URL.Scheme == "rtmp" { + dial = &net.Dialer{} + } else { + dial = &ctls.Dialer{Config: c.TLSConfig} + } + + var err error + c.nconn, err = dial.DialContext(ctx, "tcp", c.URL.Host) + if err != nil { + return err + } + + closerDone := make(chan struct{}) + defer func() { <-closerDone }() + + closerTerminate := make(chan struct{}) + defer close(closerTerminate) + + nc := c.nconn + go func() { + defer close(closerDone) + + select { + case <-closerTerminate: + case <-ctx.Done(): + nc.Close() + } + }() + + err = c.initialize3() + if err != nil { + c.nconn.Close() + return err + } + + return nil +} + +func (c *Client) initialize3() error { + c.bc = bytecounter.NewReadWriter(c.nconn) + + _, _, err := handshake.DoClient(c.bc, false, false) + if err != nil { + return err + } + + c.mrw = message.NewReadWriter(c.bc, c.bc, false) + + err = c.mrw.Write(&message.SetWindowAckSize{ + Value: 2500000, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + if err != nil { + return err + } + + cleanURL := &url.URL{ + Scheme: c.URL.Scheme, + Opaque: c.URL.Opaque, + Host: c.URL.Host, + Path: c.URL.Path, + RawPath: c.URL.RawPath, + OmitHost: c.URL.OmitHost, + ForceQuery: c.URL.ForceQuery, + RawQuery: c.URL.RawQuery, + Fragment: c.URL.Fragment, + RawFragment: c.URL.RawFragment, + } + app, streamKey := splitPath(cleanURL) + tcURL := getTcURL(cleanURL) + + switch c.authState { + case 1: + user := c.URL.User.Username() + + app += "?authmod=adobe&user=" + user + tcURL += "?authmod=adobe&user=" + user + + case 2: + user := c.URL.User.Username() + pass, _ := c.URL.User.Password() + + clientChallenge := strings.ReplaceAll(uuid.New().String(), "-", "") + response := authResponse(user, pass, c.authSalt, "", c.authChallenge, clientChallenge) + + app += fmt.Sprintf("?authmod=adobe&user=myuser&challenge=%s&response=%s", clientChallenge, response) + tcURL += fmt.Sprintf("?authmod=adobe&user=myuser&challenge=%s&response=%s", clientChallenge, response) + } + + connectArg := amf0.Object{ + {Key: "app", Value: app}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: tcURL}, + } + + if !c.Publish { + connectArg = append(connectArg, + amf0.ObjectEntry{Key: "fpad", Value: false}, + amf0.ObjectEntry{Key: "capabilities", Value: float64(15)}, + amf0.ObjectEntry{Key: "audioCodecs", Value: float64(4071)}, + amf0.ObjectEntry{Key: "videoCodecs", Value: float64(252)}, + amf0.ObjectEntry{Key: "videoFunction", Value: float64(1)}, + ) + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{connectArg}, + }) + if err != nil { + return err + } + + res, err := readCommandResult(c.mrw, 1) + if err != nil { + return err + } + + switch res.Name { + case "_result": + + case "_error": + if len(res.Arguments) < 2 { + return fmt.Errorf("bad result: %v", res) + } + + ma, ok := objectOrArray(res.Arguments[1]) + if !ok { + return fmt.Errorf("bad result: %v", res) + } + + desc, ok := ma.GetString("description") + if !ok { + return fmt.Errorf("bad result: %v", res) + } + + if desc == "code=403 need auth; authmod=adobe" { + if c.URL.User == nil { + return fmt.Errorf("credentials are required") + } + + if c.authState != 0 { + return fmt.Errorf("authentication error") + } + + return errAuth + } + + if !strings.HasPrefix(desc, "authmod=adobe ?") { + return fmt.Errorf("bad result: %v", res) + } + + desc = desc[len("authmod=adobe ?"):] + vals := queryDecode(desc) + + reason := vals["reason"] + c.authSalt = vals["salt"] + c.authChallenge = vals["challenge"] + + if reason != "needauth" || c.authSalt == "" || c.authChallenge == "" { + return fmt.Errorf("bad result: %v", res) + } + + if c.authState != 1 { + return fmt.Errorf("authentication error") + } + + return errAuth + + default: + return fmt.Errorf("bad result: %v", res) + } + + if !c.Publish { + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + }, + }) + if err != nil { + return err + } + + res, err = readCommandResult(c.mrw, 2) + if err != nil { + return err + } + + if res.Name != "_result" || !resultIsOK2(res) { + return fmt.Errorf("bad result: %v", res) + } + + err = c.mrw.Write(&message.UserControlSetBufferLength{ + BufferLength: 0x64, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "play", + CommandID: 3, + Arguments: []interface{}{ + nil, + streamKey, + }, + }) + if err != nil { + return err + } + + res, err = readCommandResult(c.mrw, 3) + if err != nil { + return err + } + + if res.Name != "onStatus" || !resultIsOK1(res) { + return fmt.Errorf("bad result: %v", res) + } + } else { + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "releaseStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + streamKey, + }, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "FCPublish", + CommandID: 3, + Arguments: []interface{}{ + nil, + streamKey, + }, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 4, + Arguments: []interface{}{ + nil, + }, + }) + if err != nil { + return err + } + + res, err = readCommandResult(c.mrw, 4) + if err != nil { + return err + } + + if res.Name != "_result" || !resultIsOK2(res) { + return fmt.Errorf("bad result: %v", res) + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "publish", + CommandID: 5, + Arguments: []interface{}{ + nil, + streamKey, + app, + }, + }) + if err != nil { + return err + } + + res, err = readCommandResult(c.mrw, 5) + if err != nil { + return err + } + + if res.Name != "onStatus" || !resultIsOK1(res) { + return fmt.Errorf("bad result: %v", res) + } + } + + return nil +} + +// Close closes the connection. +func (c *Client) Close() { + c.nconn.Close() +} + +// NetConn returns the underlying net.Conn. +func (c *Client) NetConn() net.Conn { + return c.nconn +} + +// BytesReceived returns the number of bytes received. +func (c *Client) BytesReceived() uint64 { + return c.bc.Reader.Count() +} + +// BytesSent returns the number of bytes sent. +func (c *Client) BytesSent() uint64 { + return c.bc.Writer.Count() +} + +// Read reads a message. +func (c *Client) Read() (message.Message, error) { + return c.mrw.Read() +} + +// Write writes a message. +func (c *Client) Write(msg message.Message) error { + return c.mrw.Write(msg) +} diff --git a/internal/protocols/rtmp/client_test.go b/internal/protocols/rtmp/client_test.go new file mode 100644 index 00000000..f3380c35 --- /dev/null +++ b/internal/protocols/rtmp/client_test.go @@ -0,0 +1,421 @@ +package rtmp + +import ( + "context" + "net" + "net/url" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" +) + +func TestClient(t *testing.T) { + for _, ca := range []string{ + "auth", + "read", + "read nginx rtmp", + "publish", + } { + t.Run(ca, func(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:9121") + require.NoError(t, err) + defer ln.Close() + + done := make(chan struct{}) + authState := 0 + + go func() { + for { + conn, err2 := ln.Accept() + require.NoError(t, err2) + defer conn.Close() + bc := bytecounter.NewReadWriter(conn) + + _, _, err2 = handshake.DoServer(bc, false) + require.NoError(t, err2) + + mrw := message.NewReadWriter(bc, bc, true) + + msg, err2 := mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.SetWindowAckSize{ + Value: 2500000, + }, msg) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }, msg) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.SetChunkSize{ + Value: 65536, + }, msg) + + switch ca { + case "auth": + msg, err2 = mrw.Read() + require.NoError(t, err2) + + switch authState { + case 0: + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }, msg) + + case 1: + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream?authmod=adobe&user=myuser"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream?authmod=adobe&user=myuser"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }, msg) + + case 2: + app, _ := msg.(*message.CommandAMF0).Arguments[0].(amf0.Object).GetString("app") + query := queryDecode(app[len("stream?"):]) + clientChallenge := query["challenge"] + response := authResponse("myuser", "mypass", "salt123", "", "server456challenge", clientChallenge) + + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + { + Key: "app", + Value: "stream?authmod=adobe&user=myuser&challenge=" + + clientChallenge + "&response=" + response, + }, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + { + Key: "tcUrl", + Value: "rtmp://127.0.0.1:9121/stream?authmod=adobe&user=myuser&challenge=" + + clientChallenge + "&response=" + response, + }, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }, msg) + } + + case "read", "read nginx rtmp": + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }, msg) + + case "publish": + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream"}, + }, + }, + }, msg) + } + + if ca == "auth" { + switch authState { + case 0: + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_error", + CommandID: 1, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + {Key: "description", Value: "code=403 need auth; authmod=adobe"}, + }, + }, + }) + require.NoError(t, err2) + + authState++ + continue + + case 1: + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_error", + CommandID: 1, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + { + Key: "description", + Value: "authmod=adobe ?reason=needauth&user=myuser&salt=salt123&challenge=server456challenge", + }, + }, + }, + }) + require.NoError(t, err2) + + authState++ + continue + } + } + + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: float64(0)}, + }, + }, + }) + require.NoError(t, err2) + + switch ca { + case "auth", "read", "read nginx rtmp": + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + }, + }, msg) + + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 2, + Arguments: []interface{}{ + nil, + float64(1), + }, + }) + require.NoError(t, err2) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.UserControlSetBufferLength{ + BufferLength: 0x64, + }, msg) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "play", + CommandID: 3, + Arguments: []interface{}{ + nil, + "", + }, + }, msg) + + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: func() int { + if ca == "read nginx rtmp" { + return 0 + } + return 3 + }(), + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Play.Reset"}, + {Key: "description", Value: "play reset"}, + }, + }, + }) + require.NoError(t, err2) + + case "publish": + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "releaseStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + "", + }, + }, msg) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "FCPublish", + CommandID: 3, + Arguments: []interface{}{ + nil, + "", + }, + }, msg) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 4, + Arguments: []interface{}{ + nil, + }, + }, msg) + + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 4, + Arguments: []interface{}{ + nil, + float64(1), + }, + }) + require.NoError(t, err2) + + msg, err2 = mrw.Read() + require.NoError(t, err2) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "publish", + CommandID: 5, + Arguments: []interface{}{ + nil, + "", + "stream", + }, + }, msg) + + err2 = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: 5, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Publish.Start"}, + {Key: "description", Value: "publish start"}, + }, + }, + }) + require.NoError(t, err2) + } + + close(done) + break + } + }() + + var rawURL string + + if ca == "auth" { + rawURL = "rtmp://myuser:mypass@127.0.0.1:9121/stream" + } else { + rawURL = "rtmp://127.0.0.1:9121/stream" + } + + u, err := url.Parse(rawURL) + require.NoError(t, err) + + conn := &Client{ + URL: u, + Publish: (ca == "publish"), + } + err = conn.Initialize(context.Background()) + require.NoError(t, err) + defer conn.Close() + + switch ca { + case "read", "read nginx rtmp": + require.Equal(t, uint64(3421), conn.BytesReceived()) + require.Equal(t, uint64(3409), conn.BytesSent()) + + case "publish": + require.Equal(t, uint64(3427), conn.BytesReceived()) + require.Equal(t, uint64(0xd27), conn.BytesSent()) + } + + <-done + }) + } +} diff --git a/internal/protocols/rtmp/conn.go b/internal/protocols/rtmp/conn.go index 5b34292d..3098cb6a 100644 --- a/internal/protocols/rtmp/conn.go +++ b/internal/protocols/rtmp/conn.go @@ -1,635 +1,46 @@ -// Package rtmp provides RTMP connectivity. package rtmp import ( - "fmt" "io" - "net/url" - "strings" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" ) -func resultIsOK1(res *message.CommandAMF0) bool { - if len(res.Arguments) < 2 { - return false - } - - var ma amf0.Object - switch pl := res.Arguments[1].(type) { - case amf0.Object: - ma = pl - - case amf0.ECMAArray: - ma = amf0.Object(pl) - - default: - return false - } - - v, ok := ma.Get("level") - if !ok { - return false - } - - return v == "status" +// Conn is implemented by Client and ServerConn. +type Conn interface { + BytesReceived() uint64 + BytesSent() uint64 + Read() (message.Message, error) + Write(msg message.Message) error } -func resultIsOK2(res *message.CommandAMF0) bool { - if len(res.Arguments) < 2 { - return false - } - - v, ok := res.Arguments[1].(float64) - if !ok { - return false - } - - return v == 1 -} - -func splitPath(u *url.URL) (app, stream string) { - nu := *u - nu.ForceQuery = false - - pathsegs := strings.Split(nu.RequestURI(), "/") - if len(pathsegs) == 2 { - app = pathsegs[1] - } - if len(pathsegs) == 3 { - app = pathsegs[1] - stream = pathsegs[2] - } - if len(pathsegs) > 3 { - app = strings.Join(pathsegs[1:3], "/") - stream = strings.Join(pathsegs[3:], "/") - } - return -} - -func getTcURL(u *url.URL) string { - app, _ := splitPath(u) - nu, _ := url.Parse(u.String()) // perform a deep copy - nu.RawQuery = "" - nu.Path = "/" - return nu.String() + app -} - -func createURL(tcURL string, app string, play string) (*url.URL, error) { - u, err := url.ParseRequestURI("/" + app + "/" + play) - if err != nil { - return nil, err - } - - tu, err := url.Parse(tcURL) - if err != nil { - return nil, err - } - - if tu.Host == "" { - return nil, fmt.Errorf("invalid host") - } - u.Host = tu.Host - - if tu.Scheme == "" { - return nil, fmt.Errorf("invalid scheme") - } - u.Scheme = tu.Scheme - - return u, nil -} - -func readCommand(mrw *message.ReadWriter) (*message.CommandAMF0, error) { - for { - msg, err := mrw.Read() - if err != nil { - return nil, err - } - - if cmd, ok := msg.(*message.CommandAMF0); ok { - return cmd, nil - } - } -} - -func readCommandResult( - mrw *message.ReadWriter, - commandID int, - commandName string, - isValid func(*message.CommandAMF0) bool, -) error { - for { - msg, err := mrw.Read() - if err != nil { - return err - } - - if cmd, ok := msg.(*message.CommandAMF0); ok { - if (cmd.CommandID == commandID || cmd.CommandID == 0) && cmd.Name == commandName { - if !isValid(cmd) { - return fmt.Errorf("server refused connect request") - } - - return nil - } - } - } -} - -// Conn is a RTMP connection. -type Conn struct { - RW io.ReadWriter - Client bool - URL *url.URL - Publish bool - skipHandshake bool +type dummyConn struct { + rw io.ReadWriter bc *bytecounter.ReadWriter mrw *message.ReadWriter } -// Initialize initializes Conn. -func (c *Conn) Initialize() error { - c.bc = bytecounter.NewReadWriter(c.RW) - - if !c.skipHandshake { - if c.Client { - if c.URL == nil { - return fmt.Errorf("URL must be specified in client mode") - } - - err := c.initializeClient() - if err != nil { - return err - } - } else { - if c.URL != nil { - return fmt.Errorf("URL must be empty in server mode") - } - - var err error - c.URL, c.Publish, err = c.initializeServer() - if err != nil { - return err - } - } - } else { - c.mrw = message.NewReadWriter(c.bc, c.bc, false) - } - - return nil -} - -func (c *Conn) initializeClient() error { - connectpath, actionpath := splitPath(c.URL) - - _, _, err := handshake.DoClient(c.bc, false, false) - if err != nil { - return err - } - +func (c *dummyConn) initialize() { + c.bc = bytecounter.NewReadWriter(c.rw) c.mrw = message.NewReadWriter(c.bc, c.bc, false) - - err = c.mrw.Write(&message.SetWindowAckSize{ - Value: 2500000, - }) - if err != nil { - return err - } - - err = c.mrw.Write(&message.SetPeerBandwidth{ - Value: 2500000, - Type: 2, - }) - if err != nil { - return err - } - - err = c.mrw.Write(&message.SetChunkSize{ - Value: 65536, - }) - if err != nil { - return err - } - - connectArg := amf0.Object{ - {Key: "app", Value: connectpath}, - {Key: "flashVer", Value: "LNX 9,0,124,2"}, - {Key: "tcUrl", Value: getTcURL(c.URL)}, - } - - if !c.Publish { - connectArg = append(connectArg, - amf0.ObjectEntry{Key: "fpad", Value: false}, - amf0.ObjectEntry{Key: "capabilities", Value: float64(15)}, - amf0.ObjectEntry{Key: "audioCodecs", Value: float64(4071)}, - amf0.ObjectEntry{Key: "videoCodecs", Value: float64(252)}, - amf0.ObjectEntry{Key: "videoFunction", Value: float64(1)}, - ) - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "connect", - CommandID: 1, - Arguments: []interface{}{connectArg}, - }) - if err != nil { - return err - } - - err = readCommandResult(c.mrw, 1, "_result", resultIsOK1) - if err != nil { - return err - } - - if !c.Publish { - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - }, - }) - if err != nil { - return err - } - - err = readCommandResult(c.mrw, 2, "_result", resultIsOK2) - if err != nil { - return err - } - - err = c.mrw.Write(&message.UserControlSetBufferLength{ - BufferLength: 0x64, - }) - if err != nil { - return err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "play", - CommandID: 3, - Arguments: []interface{}{ - nil, - actionpath, - }, - }) - if err != nil { - return err - } - - return readCommandResult(c.mrw, 3, "onStatus", resultIsOK1) - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "releaseStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - actionpath, - }, - }) - if err != nil { - return err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "FCPublish", - CommandID: 3, - Arguments: []interface{}{ - nil, - actionpath, - }, - }) - if err != nil { - return err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 4, - Arguments: []interface{}{ - nil, - }, - }) - if err != nil { - return err - } - - err = readCommandResult(c.mrw, 4, "_result", resultIsOK2) - if err != nil { - return err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "publish", - CommandID: 5, - Arguments: []interface{}{ - nil, - actionpath, - connectpath, - }, - }) - if err != nil { - return err - } - - return readCommandResult(c.mrw, 5, "onStatus", resultIsOK1) -} - -func (c *Conn) initializeServer() (*url.URL, bool, error) { - keyIn, keyOut, err := handshake.DoServer(c.bc, false) - if err != nil { - return nil, false, err - } - - var rw io.ReadWriter - if keyIn != nil { - rw, err = newRC4ReadWriter(c.bc, keyIn, keyOut) - if err != nil { - return nil, false, err - } - } else { - rw = c.bc - } - - c.mrw = message.NewReadWriter(rw, c.bc, false) - - cmd, err := readCommand(c.mrw) - if err != nil { - return nil, false, err - } - - if cmd.Name != "connect" { - return nil, false, fmt.Errorf("unexpected command: %+v", cmd) - } - - if len(cmd.Arguments) < 1 { - return nil, false, fmt.Errorf("invalid connect command: %+v", cmd) - } - - var ma amf0.Object - switch pl := cmd.Arguments[0].(type) { - case amf0.Object: - ma = pl - - case amf0.ECMAArray: - ma = amf0.Object(pl) - - default: - return nil, false, fmt.Errorf("invalid connect command: %+v", cmd) - } - - connectpath, ok := ma.GetString("app") - if !ok { - return nil, false, fmt.Errorf("invalid connect command: %+v", cmd) - } - - tcURL, ok := ma.GetString("tcUrl") - if !ok { - tcURL, ok = ma.GetString("tcurl") - if !ok { - return nil, false, fmt.Errorf("invalid connect command: %+v", cmd) - } - } - - tcURL = strings.Trim(tcURL, "'") - - err = c.mrw.Write(&message.SetWindowAckSize{ - Value: 2500000, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.SetPeerBandwidth{ - Value: 2500000, - Type: 2, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.SetChunkSize{ - Value: 65536, - }) - if err != nil { - return nil, false, err - } - - oe, _ := ma.GetFloat64("objectEncoding") - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: cmd.ChunkStreamID, - Name: "_result", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - amf0.Object{ - {Key: "fmsVer", Value: "LNX 9,0,124,2"}, - {Key: "capabilities", Value: float64(31)}, - }, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetConnection.Connect.Success"}, - {Key: "description", Value: "Connection succeeded."}, - {Key: "objectEncoding", Value: oe}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - for { - cmd, err := readCommand(c.mrw) - if err != nil { - return nil, false, err - } - - switch cmd.Name { - case "createStream": - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: cmd.ChunkStreamID, - Name: "_result", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - nil, - float64(1), - }, - }) - if err != nil { - return nil, false, err - } - - case "play": - if len(cmd.Arguments) < 2 { - return nil, false, fmt.Errorf("invalid play command arguments") - } - - actionpath, ok := cmd.Arguments[1].(string) - if !ok { - return nil, false, fmt.Errorf("invalid play command arguments") - } - - u, err := createURL(tcURL, connectpath, actionpath) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.UserControlStreamIsRecorded{ - StreamID: 1, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.UserControlStreamBegin{ - StreamID: 1, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Play.Reset"}, - {Key: "description", Value: "play reset"}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Play.Start"}, - {Key: "description", Value: "play start"}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Data.Start"}, - {Key: "description", Value: "data start"}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: cmd.CommandID, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Play.PublishNotify"}, - {Key: "description", Value: "publish notify"}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - return u, false, nil - - case "publish": - if len(cmd.Arguments) < 2 { - return nil, false, fmt.Errorf("invalid publish command arguments") - } - - actionpath, ok := cmd.Arguments[1].(string) - if !ok { - return nil, false, fmt.Errorf("invalid publish command arguments") - } - - u, err := createURL(tcURL, connectpath, actionpath) - if err != nil { - return nil, false, err - } - - err = c.mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - Name: "onStatus", - CommandID: cmd.CommandID, - MessageStreamID: 0x1000000, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Publish.Start"}, - {Key: "description", Value: "publish start"}, - }, - }, - }) - if err != nil { - return nil, false, err - } - - return u, true, nil - } - } } // BytesReceived returns the number of bytes received. -func (c *Conn) BytesReceived() uint64 { +func (c *dummyConn) BytesReceived() uint64 { return c.bc.Reader.Count() } // BytesSent returns the number of bytes sent. -func (c *Conn) BytesSent() uint64 { +func (c *dummyConn) BytesSent() uint64 { return c.bc.Writer.Count() } -// Read reads a message. -func (c *Conn) Read() (message.Message, error) { +func (c *dummyConn) Read() (message.Message, error) { return c.mrw.Read() } -// Write writes a message. -func (c *Conn) Write(msg message.Message) error { +func (c *dummyConn) Write(msg message.Message) error { return c.mrw.Write(msg) } diff --git a/internal/protocols/rtmp/conn_test.go b/internal/protocols/rtmp/conn_test.go deleted file mode 100644 index fa1373ec..00000000 --- a/internal/protocols/rtmp/conn_test.go +++ /dev/null @@ -1,542 +0,0 @@ -package rtmp - -import ( - "bytes" - "net" - "net/url" - "testing" - - "github.com/stretchr/testify/require" - - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" -) - -func TestNewClientConn(t *testing.T) { - for _, ca := range []string{ - "read", - "read nginx rtmp", - "publish", - } { - t.Run(ca, func(t *testing.T) { - ln, err := net.Listen("tcp", "127.0.0.1:9121") - require.NoError(t, err) - defer ln.Close() - - done := make(chan struct{}) - - go func() { - conn, err2 := ln.Accept() - require.NoError(t, err2) - defer conn.Close() - bc := bytecounter.NewReadWriter(conn) - - _, _, err2 = handshake.DoServer(bc, false) - require.NoError(t, err2) - - mrw := message.NewReadWriter(bc, bc, true) - - msg, err2 := mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.SetWindowAckSize{ - Value: 2500000, - }, msg) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.SetPeerBandwidth{ - Value: 2500000, - Type: 2, - }, msg) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.SetChunkSize{ - Value: 65536, - }, msg) - - if ca != "publish" { - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "connect", - CommandID: 1, - Arguments: []interface{}{ - amf0.Object{ - {Key: "app", Value: "stream"}, - {Key: "flashVer", Value: "LNX 9,0,124,2"}, - {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream"}, - {Key: "fpad", Value: false}, - {Key: "capabilities", Value: float64(15)}, - {Key: "audioCodecs", Value: float64(4071)}, - {Key: "videoCodecs", Value: float64(252)}, - {Key: "videoFunction", Value: float64(1)}, - }, - }, - }, msg) - } else { - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "connect", - CommandID: 1, - Arguments: []interface{}{ - amf0.Object{ - {Key: "app", Value: "stream"}, - {Key: "flashVer", Value: "LNX 9,0,124,2"}, - {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream"}, - }, - }, - }, msg) - } - - err2 = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 1, - Arguments: []interface{}{ - amf0.Object{ - {Key: "fmsVer", Value: "LNX 9,0,124,2"}, - {Key: "capabilities", Value: float64(31)}, - }, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetConnection.Connect.Success"}, - {Key: "description", Value: "Connection succeeded."}, - {Key: "objectEncoding", Value: float64(0)}, - }, - }, - }) - require.NoError(t, err2) - - switch ca { - case "read", "read nginx rtmp": - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - }, - }, msg) - - err2 = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 2, - Arguments: []interface{}{ - nil, - float64(1), - }, - }) - require.NoError(t, err2) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.UserControlSetBufferLength{ - BufferLength: 0x64, - }, msg) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "play", - CommandID: 3, - Arguments: []interface{}{ - nil, - "", - }, - }, msg) - - err2 = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: func() int { - if ca == "read nginx rtmp" { - return 0 - } - return 3 - }(), - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Play.Reset"}, - {Key: "description", Value: "play reset"}, - }, - }, - }) - require.NoError(t, err2) - - case "publish": - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "releaseStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - "", - }, - }, msg) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "FCPublish", - CommandID: 3, - Arguments: []interface{}{ - nil, - "", - }, - }, msg) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 4, - Arguments: []interface{}{ - nil, - }, - }, msg) - - err2 = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 4, - Arguments: []interface{}{ - nil, - float64(1), - }, - }) - require.NoError(t, err2) - - msg, err2 = mrw.Read() - require.NoError(t, err2) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "publish", - CommandID: 5, - Arguments: []interface{}{ - nil, - "", - "stream", - }, - }, msg) - - err2 = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 5, - MessageStreamID: 0x1000000, - Name: "onStatus", - CommandID: 5, - Arguments: []interface{}{ - nil, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetStream.Publish.Start"}, - {Key: "description", Value: "publish start"}, - }, - }, - }) - require.NoError(t, err2) - } - - close(done) - }() - - u, err := url.Parse("rtmp://127.0.0.1:9121/stream") - require.NoError(t, err) - - nconn, err := net.Dial("tcp", u.Host) - require.NoError(t, err) - defer nconn.Close() - - conn := &Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: ca == "publish", - } - err = conn.Initialize() - require.NoError(t, err) - - switch ca { - case "read", "read nginx rtmp": - require.Equal(t, uint64(3421), conn.BytesReceived()) - require.Equal(t, uint64(3409), conn.BytesSent()) - - case "publish": - require.Equal(t, uint64(3427), conn.BytesReceived()) - require.Equal(t, uint64(0xd27), conn.BytesSent()) - } - - <-done - }) - } -} - -func TestNewServerConn(t *testing.T) { - for _, ca := range []string{ - "read", - "publish", - "publish neko", - } { - t.Run(ca, func(t *testing.T) { - ln, err := net.Listen("tcp", "127.0.0.1:9121") - require.NoError(t, err) - defer ln.Close() - - done := make(chan struct{}) - - go func() { - nconn, err2 := ln.Accept() - require.NoError(t, err2) - defer nconn.Close() - - conn := &Conn{ - RW: nconn, - Client: false, - } - err2 = conn.Initialize() - require.NoError(t, err2) - - require.Equal(t, &url.URL{ - Scheme: "rtmp", - Host: "127.0.0.1:9121", - Path: "//stream/", - }, conn.URL) - require.Equal(t, ca == "publish" || ca == "publish neko", conn.Publish) - - close(done) - }() - - conn, err := net.Dial("tcp", "127.0.0.1:9121") - require.NoError(t, err) - defer conn.Close() - bc := bytecounter.NewReadWriter(conn) - - _, _, err = handshake.DoClient(bc, false, false) - require.NoError(t, err) - - mrw := message.NewReadWriter(bc, bc, true) - - tcURL := "rtmp://127.0.0.1:9121/stream" - if ca == "publish neko" { - tcURL = "'rtmp://127.0.0.1:9121/stream" - } - - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "connect", - CommandID: 1, - Arguments: []interface{}{ - amf0.Object{ - {Key: "app", Value: "/stream"}, - {Key: "flashVer", Value: "LNX 9,0,124,2"}, - {Key: "tcUrl", Value: tcURL}, - {Key: "fpad", Value: false}, - {Key: "capabilities", Value: float64(15)}, - {Key: "audioCodecs", Value: float64(4071)}, - {Key: "videoCodecs", Value: float64(252)}, - {Key: "videoFunction", Value: float64(1)}, - }, - }, - }) - require.NoError(t, err) - - msg, err := mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.SetWindowAckSize{ - Value: 2500000, - }, msg) - - msg, err = mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.SetPeerBandwidth{ - Value: 2500000, - Type: 2, - }, msg) - - msg, err = mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.SetChunkSize{ - Value: 65536, - }, msg) - - msg, err = mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 1, - Arguments: []interface{}{ - amf0.Object{ - {Key: "fmsVer", Value: "LNX 9,0,124,2"}, - {Key: "capabilities", Value: float64(31)}, - }, - amf0.Object{ - {Key: "level", Value: "status"}, - {Key: "code", Value: "NetConnection.Connect.Success"}, - {Key: "description", Value: "Connection succeeded."}, - {Key: "objectEncoding", Value: float64(0)}, - }, - }, - }, msg) - - err = mrw.Write(&message.SetChunkSize{ - Value: 65536, - }) - require.NoError(t, err) - - if ca == "read" { - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - }, - }) - require.NoError(t, err) - - msg, err = mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 2, - Arguments: []interface{}{ - nil, - float64(1), - }, - }, msg) - - err = mrw.Write(&message.UserControlSetBufferLength{ - BufferLength: 0x64, - }) - require.NoError(t, err) - - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "play", - CommandID: 0, - Arguments: []interface{}{ - nil, - "", - }, - }) - require.NoError(t, err) - } else { - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "releaseStream", - CommandID: 2, - Arguments: []interface{}{ - nil, - "", - }, - }) - require.NoError(t, err) - - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "FCPublish", - CommandID: 3, - Arguments: []interface{}{ - nil, - "", - }, - }) - require.NoError(t, err) - - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 3, - Name: "createStream", - CommandID: 4, - Arguments: []interface{}{ - nil, - }, - }) - require.NoError(t, err) - - msg, err = mrw.Read() - require.NoError(t, err) - require.Equal(t, &message.CommandAMF0{ - ChunkStreamID: 3, - Name: "_result", - CommandID: 4, - Arguments: []interface{}{ - nil, - float64(1), - }, - }, msg) - - err = mrw.Write(&message.CommandAMF0{ - ChunkStreamID: 4, - MessageStreamID: 0x1000000, - Name: "publish", - CommandID: 5, - Arguments: []interface{}{ - nil, - "", - "stream", - }, - }) - require.NoError(t, err) - } - - <-done - }) - } -} - -func BenchmarkRead(b *testing.B) { - var buf bytes.Buffer - - for n := 0; n < b.N; n++ { - buf.Write([]byte{ - 7, 0, 0, 23, 0, 0, 98, 8, - 0, 0, 0, 64, 175, 1, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, 1, 2, - 3, 4, 1, 2, 3, 4, - }) - } - - conn := &Conn{ - RW: &buf, - skipHandshake: true, - } - err := conn.Initialize() - if err != nil { - panic(err) - } - - for n := 0; n < b.N; n++ { - conn.Read() //nolint:errcheck - } -} diff --git a/internal/protocols/rtmp/from_stream.go b/internal/protocols/rtmp/from_stream.go index f4b3149a..bcd5cdbf 100644 --- a/internal/protocols/rtmp/from_stream.go +++ b/internal/protocols/rtmp/from_stream.go @@ -187,7 +187,7 @@ func setupAudio( func FromStream( str *stream.Stream, reader stream.Reader, - conn *Conn, + conn Conn, nconn net.Conn, writeTimeout time.Duration, ) error { diff --git a/internal/protocols/rtmp/from_stream_test.go b/internal/protocols/rtmp/from_stream_test.go index ca440f58..14240877 100644 --- a/internal/protocols/rtmp/from_stream_test.go +++ b/internal/protocols/rtmp/from_stream_test.go @@ -8,8 +8,6 @@ import ( "github.com/bluenviron/gortsplib/v4/pkg/description" "github.com/bluenviron/gortsplib/v4/pkg/format" "github.com/bluenviron/mediamtx/internal/logger" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" - "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" "github.com/bluenviron/mediamtx/internal/stream" "github.com/bluenviron/mediamtx/internal/test" "github.com/stretchr/testify/require" @@ -75,10 +73,12 @@ func TestFromStreamSkipUnsupportedTracks(t *testing.T) { }) var buf bytes.Buffer - bc := bytecounter.NewReadWriter(&buf) - conn := &Conn{mrw: message.NewReadWriter(&buf, bc, false)} + c := &dummyConn{ + rw: &buf, + } + c.initialize() - err = FromStream(strm, l, conn, nil, 0) + err = FromStream(strm, l, c, nil, 0) require.NoError(t, err) defer strm.RemoveReader(l) diff --git a/internal/protocols/rtmp/reader.go b/internal/protocols/rtmp/reader.go index d8ff52c9..5aa85be2 100644 --- a/internal/protocols/rtmp/reader.go +++ b/internal/protocols/rtmp/reader.go @@ -270,9 +270,9 @@ func sortedKeys(m map[uint8]format.Format) []int { return ret } -// Reader is a wrapper around Conn that provides utilities to demux incoming data. +// Reader provides functions to read incoming data. type Reader struct { - Conn *Conn + Conn Conn videoTracks map[uint8]format.Format audioTracks map[uint8]format.Format @@ -280,7 +280,7 @@ type Reader struct { onAudioData map[uint8]func(message.Message) error } -// Initialize initializes a reader. +// Initialize initializes Reader. func (r *Reader) Initialize() error { var err error r.videoTracks, r.audioTracks, err = r.readTracks() diff --git a/internal/protocols/rtmp/reader_test.go b/internal/protocols/rtmp/reader_test.go index f53ffb2d..e6ab2957 100644 --- a/internal/protocols/rtmp/reader_test.go +++ b/internal/protocols/rtmp/reader_test.go @@ -1626,16 +1626,14 @@ func TestReadTracks(t *testing.T) { mrw := message.NewReadWriter(bc, bc, true) for _, msg := range ca.messages { - err := mrw.Write(msg) + err = mrw.Write(msg) require.NoError(t, err) } - c := &Conn{ - RW: &buf, - skipHandshake: true, + c := &dummyConn{ + rw: &buf, } - err := c.Initialize() - require.NoError(t, err) + c.initialize() r := &Reader{ Conn: c, diff --git a/internal/protocols/rtmp/server_conn.go b/internal/protocols/rtmp/server_conn.go new file mode 100644 index 00000000..3ca9e456 --- /dev/null +++ b/internal/protocols/rtmp/server_conn.go @@ -0,0 +1,509 @@ +package rtmp + +import ( + "crypto/md5" + "encoding/base64" + "fmt" + "io" + "net/url" + "strings" + + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" +) + +const ( + serverSalt = "testsalt" + serverChallenge = "testchallenge" +) + +func queryDecode(enc string) map[string]string { + // do not use url.ParseQuery since values are not URL-encoded + vals := make(map[string]string) + + for _, kv := range strings.Split(enc, "&") { + tmp := strings.SplitN(kv, "=", 2) + if len(tmp) == 2 { + vals[tmp[0]] = tmp[1] + } + } + + return vals +} + +func queryEncode(dec map[string]string) string { + tmp := make([]string, len(dec)) + i := 0 + + for k, v := range dec { + tmp[i] = k + "=" + v + i++ + } + + return strings.Join(tmp, "&") +} + +func authResponse(user, pass, salt, opaque, challenge, challenge2 string) string { + h := md5.New() + h.Write([]byte(user)) + h.Write([]byte(salt)) + h.Write([]byte(pass)) + str := base64.StdEncoding.EncodeToString(h.Sum(nil)) + + h = md5.New() + h.Write([]byte(str)) + if opaque != "" { + h.Write([]byte(opaque)) + } else { + h.Write([]byte(challenge)) + } + h.Write([]byte(challenge2)) + return base64.StdEncoding.EncodeToString(h.Sum(nil)) +} + +func buildURL(tcURL string, app string, streamKey string) (*url.URL, error) { + raw := "/" + app + if streamKey != "" { + raw += "/" + streamKey + } + + u, err := url.ParseRequestURI(raw) + if err != nil { + return nil, err + } + + tu, err := url.Parse(tcURL) + if err != nil { + return nil, err + } + + if tu.Host == "" { + return nil, fmt.Errorf("invalid host") + } + u.Host = tu.Host + + if tu.Scheme == "" { + return nil, fmt.Errorf("invalid scheme") + } + u.Scheme = tu.Scheme + + return u, nil +} + +func objectOrArray(in interface{}) (amf0.Object, bool) { + switch o := in.(type) { + case amf0.Object: + return o, true + + case amf0.ECMAArray: + return amf0.Object(o), true + + default: + return nil, false + } +} + +// ServerConn is a server-side RTMP connection. +type ServerConn struct { + RW io.ReadWriter + + // filled by Initialize + connectCmd *message.CommandAMF0 + connectObject amf0.Object + app string + tcURL string + + // filled by Accept + URL *url.URL + Publish bool + + bc *bytecounter.ReadWriter + mrw *message.ReadWriter +} + +// Initialize initializes ServerConn. +func (c *ServerConn) Initialize() error { + c.bc = bytecounter.NewReadWriter(c.RW) + + keyIn, keyOut, err := handshake.DoServer(c.bc, false) + if err != nil { + return err + } + + var rw io.ReadWriter + if keyIn != nil { + rw, err = newRC4ReadWriter(c.bc, keyIn, keyOut) + if err != nil { + return err + } + } else { + rw = c.bc + } + + c.mrw = message.NewReadWriter(rw, c.bc, false) + + c.connectCmd, err = readCommand(c.mrw) + if err != nil { + return err + } + + if c.connectCmd.Name != "connect" { + return fmt.Errorf("unexpected command: %+v", c.connectCmd) + } + + if len(c.connectCmd.Arguments) < 1 { + return fmt.Errorf("invalid connect command: %+v", c.connectCmd) + } + + var ok bool + c.connectObject, ok = objectOrArray(c.connectCmd.Arguments[0]) + if !ok { + return fmt.Errorf("invalid connect command: %+v", c.connectCmd) + } + + c.app, ok = c.connectObject.GetString("app") + if !ok { + return fmt.Errorf("invalid connect command: %+v", c.connectCmd) + } + + c.tcURL, ok = c.connectObject.GetString("tcUrl") + if !ok { + c.tcURL, ok = c.connectObject.GetString("tcurl") + if !ok { + return fmt.Errorf("invalid connect command: %+v", c.connectCmd) + } + } + + c.tcURL = strings.Trim(c.tcURL, "'") + + return nil +} + +// CheckCredentials checks credentials. +func (c *ServerConn) CheckCredentials(expectedUser string, expectedPass string) error { + i := strings.Index(c.app, "?authmod=adobe") + if i < 0 { + err := c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: c.connectCmd.ChunkStreamID, + Name: "_error", + CommandID: c.connectCmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + {Key: "description", Value: "code=403 need auth; authmod=adobe"}, + }, + }, + }) + if err != nil { + return err + } + + return fmt.Errorf("need auth") + } + + authParams := c.app[i+1:] + vals := queryDecode(authParams) + + user := vals["user"] + if user == "" { + return fmt.Errorf("user not provided") + } + + clientChallenge := vals["challenge"] + response := vals["response"] + + if clientChallenge == "" || response == "" { + err := c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: c.connectCmd.ChunkStreamID, + Name: "_error", + CommandID: c.connectCmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + { + Key: "description", + Value: fmt.Sprintf("authmod=adobe ?reason=needauth&user=%s&salt=%s&challenge=%s", + user, serverSalt, serverChallenge), + }, + }, + }, + }) + if err != nil { + return err + } + + return fmt.Errorf("need auth 2") + } + + expectedResponse := authResponse(expectedUser, expectedPass, serverSalt, "", serverChallenge, clientChallenge) + if expectedResponse != response { + err := c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: c.connectCmd.ChunkStreamID, + Name: "_error", + CommandID: c.connectCmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + {Key: "description", Value: "authmod=adobe ?reason=authfailed"}, + }, + }, + }) + if err != nil { + return err + } + + return fmt.Errorf("authentication failed") + } + + // remove auth parameters from app + c.app = c.app[:i] + delete(vals, "authmod") + delete(vals, "user") + delete(vals, "challenge") + delete(vals, "response") + q := queryEncode(vals) + if q != "" { + c.app += "?" + q + } + + return nil +} + +// Accept accepts the connection. +func (c *ServerConn) Accept() error { + err := c.mrw.Write(&message.SetWindowAckSize{ + Value: 2500000, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + if err != nil { + return err + } + + oe, _ := c.connectObject.GetFloat64("objectEncoding") + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: c.connectCmd.ChunkStreamID, + Name: "_result", + CommandID: c.connectCmd.CommandID, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: oe}, + }, + }, + }) + if err != nil { + return err + } + + for { + cmd, err := readCommand(c.mrw) + if err != nil { + return err + } + + switch cmd.Name { + case "createStream": + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: cmd.ChunkStreamID, + Name: "_result", + CommandID: cmd.CommandID, + Arguments: []interface{}{ + nil, + float64(1), + }, + }) + if err != nil { + return err + } + + case "play": + if len(cmd.Arguments) < 2 { + return fmt.Errorf("invalid play command arguments") + } + + streamKey, ok := cmd.Arguments[1].(string) + if !ok { + return fmt.Errorf("invalid play command arguments") + } + + c.URL, err = buildURL(c.tcURL, c.app, streamKey) + if err != nil { + return err + } + + err = c.mrw.Write(&message.UserControlStreamIsRecorded{ + StreamID: 1, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.UserControlStreamBegin{ + StreamID: 1, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: cmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Play.Reset"}, + {Key: "description", Value: "play reset"}, + }, + }, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: cmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Play.Start"}, + {Key: "description", Value: "play start"}, + }, + }, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: cmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Data.Start"}, + {Key: "description", Value: "data start"}, + }, + }, + }) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + MessageStreamID: 0x1000000, + Name: "onStatus", + CommandID: cmd.CommandID, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Play.PublishNotify"}, + {Key: "description", Value: "publish notify"}, + }, + }, + }) + if err != nil { + return err + } + + c.Publish = false + return nil + + case "publish": + if len(cmd.Arguments) < 2 { + return fmt.Errorf("invalid publish command arguments") + } + + streamKey, ok := cmd.Arguments[1].(string) + if !ok { + return fmt.Errorf("invalid publish command arguments") + } + + c.URL, err = buildURL(c.tcURL, c.app, streamKey) + if err != nil { + return err + } + + err = c.mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 5, + Name: "onStatus", + CommandID: cmd.CommandID, + MessageStreamID: 0x1000000, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetStream.Publish.Start"}, + {Key: "description", Value: "publish start"}, + }, + }, + }) + if err != nil { + return err + } + + c.Publish = true + return nil + } + } +} + +// BytesReceived returns the number of bytes received. +func (c *ServerConn) BytesReceived() uint64 { + return c.bc.Reader.Count() +} + +// BytesSent returns the number of bytes sent. +func (c *ServerConn) BytesSent() uint64 { + return c.bc.Writer.Count() +} + +// Read reads a message. +func (c *ServerConn) Read() (message.Message, error) { + return c.mrw.Read() +} + +// Write writes a message. +func (c *ServerConn) Write(msg message.Message) error { + return c.mrw.Write(msg) +} diff --git a/internal/protocols/rtmp/server_conn_test.go b/internal/protocols/rtmp/server_conn_test.go new file mode 100644 index 00000000..489f1b9b --- /dev/null +++ b/internal/protocols/rtmp/server_conn_test.go @@ -0,0 +1,728 @@ +package rtmp + +import ( + "fmt" + "net" + "net/url" + "testing" + + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/amf0" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/bytecounter" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/handshake" + "github.com/bluenviron/mediamtx/internal/protocols/rtmp/message" + "github.com/google/uuid" + "github.com/stretchr/testify/require" +) + +func TestServerConn(t *testing.T) { + for _, ca := range []string{ + "auth 1", + "auth 2", + "auth 3", + "read", + "publish", + } { + t.Run(ca, func(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:9121") + require.NoError(t, err) + defer ln.Close() + + done := make(chan struct{}) + + go func() { + defer close(done) + + nconn, err2 := ln.Accept() + require.NoError(t, err2) + defer nconn.Close() + + conn := &ServerConn{ + RW: nconn, + } + err2 = conn.Initialize() + require.NoError(t, err2) + + if ca == "auth 1" || ca == "auth 2" || ca == "auth 3" { + err2 = conn.CheckCredentials("myuser", "mypass") + switch ca { + case "auth 1": + require.Error(t, err2, "need auth") + return + case "auth 2": + require.Error(t, err2, "need auth 2") + return + case "auth 3": + require.NoError(t, err2) + } + } + + err2 = conn.Accept() + require.NoError(t, err2) + + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "/stream", + RawQuery: "key=val", + }, conn.URL) + require.Equal(t, (ca == "publish"), conn.Publish) + }() + + conn, err := net.Dial("tcp", "127.0.0.1:9121") + require.NoError(t, err) + defer conn.Close() + bc := bytecounter.NewReadWriter(conn) + + _, _, err = handshake.DoClient(bc, false, false) + require.NoError(t, err) + + mrw := message.NewReadWriter(bc, bc, true) + + switch ca { + case "auth 1": //nolint:dupl + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream?key=val"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream?key=val"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + var msg message.Message + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_error", + CommandID: 1, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + {Key: "description", Value: "code=403 need auth; authmod=adobe"}, + }, + }, + }, msg) + + case "auth 2": //nolint:dupl + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream?key=val?authmod=adobe&user=myuser"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream?key=val?authmod=adobe&user=myuser"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + var msg message.Message + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_error", + CommandID: 1, + Arguments: []interface{}{ + nil, + amf0.Object{ + {Key: "level", Value: "error"}, + {Key: "code", Value: "NetConnection.Connect.Rejected"}, + {Key: "description", Value: "authmod=adobe ?reason=needauth&user=myuser&salt=testsalt&challenge=testchallenge"}, + }, + }, + }, msg) + + case "auth 3": + clientChallenge := uuid.New().String() + response := authResponse("myuser", "mypass", serverSalt, "", serverChallenge, clientChallenge) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + { + Key: "app", + Value: fmt.Sprintf("stream?key=val?authmod=adobe&user=myuser&challenge=%s&response=%s", + clientChallenge, response), + }, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + { + Key: "tcUrl", + Value: fmt.Sprintf("rtmp://127.0.0.1:9121/stream?key=val?authmod=adobe&user=myuser&challenge=%s&response=%s", + clientChallenge, response), + }, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + var msg message.Message + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetWindowAckSize{ + Value: 2500000, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetChunkSize{ + Value: 65536, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: float64(0)}, + }, + }, + }, msg) + + err = mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + }, + }) + require.NoError(t, err) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 2, + Arguments: []interface{}{ + nil, + float64(1), + }, + }, msg) + + err = mrw.Write(&message.UserControlSetBufferLength{ + BufferLength: 0x64, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "play", + CommandID: 0, + Arguments: []interface{}{ + nil, + "", + }, + }) + require.NoError(t, err) + + case "read": + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream?key=val"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream?key=val"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + var msg message.Message + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetWindowAckSize{ + Value: 2500000, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetChunkSize{ + Value: 65536, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: float64(0)}, + }, + }, + }, msg) + + err = mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + }, + }) + require.NoError(t, err) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 2, + Arguments: []interface{}{ + nil, + float64(1), + }, + }, msg) + + err = mrw.Write(&message.UserControlSetBufferLength{ + BufferLength: 0x64, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "play", + CommandID: 0, + Arguments: []interface{}{ + nil, + "", + }, + }) + require.NoError(t, err) + + case "publish": + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: "stream?key=val"}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: "rtmp://127.0.0.1:9121/stream?key=val"}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + msg, err := mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetWindowAckSize{ + Value: 2500000, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetChunkSize{ + Value: 65536, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: float64(0)}, + }, + }, + }, msg) + + err = mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "releaseStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + "", + }, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "FCPublish", + CommandID: 3, + Arguments: []interface{}{ + nil, + "", + }, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 4, + Arguments: []interface{}{ + nil, + }, + }) + require.NoError(t, err) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 4, + Arguments: []interface{}{ + nil, + float64(1), + }, + }, msg) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "publish", + CommandID: 5, + Arguments: []interface{}{ + nil, + "", + "stream", + }, + }) + require.NoError(t, err) + } + + <-done + }) + } +} + +func TestServerConnPath(t *testing.T) { + for _, ca := range []string{ + "standard", + "leading slash", + "query", + "stream key", + "stream key and query", + "neko", + } { + t.Run(ca, func(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:9121") + require.NoError(t, err) + defer ln.Close() + + done := make(chan struct{}) + + go func() { + defer close(done) + + nconn, err2 := ln.Accept() + require.NoError(t, err2) + defer nconn.Close() + + conn := &ServerConn{ + RW: nconn, + } + err2 = conn.Initialize() + require.NoError(t, err2) + + err2 = conn.Accept() + require.NoError(t, err2) + + switch ca { + case "standard", "neko": + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "/stream", + }, conn.URL) + + case "leading slash": + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "//stream", + }, conn.URL) + + case "query": + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "/stream", + RawQuery: "key=val", + }, conn.URL) + + case "stream key": + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "/stream/key", + }, conn.URL) + + case "stream key and query": + require.Equal(t, &url.URL{ + Scheme: "rtmp", + Host: "127.0.0.1:9121", + Path: "/stream/key", + RawQuery: "key=val", + }, conn.URL) + } + }() + + conn, err := net.Dial("tcp", "127.0.0.1:9121") + require.NoError(t, err) + defer conn.Close() + bc := bytecounter.NewReadWriter(conn) + + _, _, err = handshake.DoClient(bc, false, false) + require.NoError(t, err) + + mrw := message.NewReadWriter(bc, bc, true) + + var app string + var tcURL string + + switch ca { + case "standard": + app = "stream" + tcURL = "rtmp://127.0.0.1:9121/stream" + + case "leading slash": + app = "/stream" + tcURL = "rtmp://127.0.0.1:9121//stream" + + case "query": + app = "stream?key=val" + tcURL = "rtmp://127.0.0.1:9121/stream?key=val" + + case "stream key": + app = "stream" + tcURL = "rtmp://127.0.0.1:9121/stream" + + case "stream key and query": + app = "stream" + tcURL = "rtmp://127.0.0.1:9121/stream" + + case "neko": + app = "stream" + tcURL = "'rtmp://127.0.0.1:9121/stream" + } + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "connect", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "app", Value: app}, + {Key: "flashVer", Value: "LNX 9,0,124,2"}, + {Key: "tcUrl", Value: tcURL}, + {Key: "fpad", Value: false}, + {Key: "capabilities", Value: float64(15)}, + {Key: "audioCodecs", Value: float64(4071)}, + {Key: "videoCodecs", Value: float64(252)}, + {Key: "videoFunction", Value: float64(1)}, + }, + }, + }) + require.NoError(t, err) + + msg, err := mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetWindowAckSize{ + Value: 2500000, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetPeerBandwidth{ + Value: 2500000, + Type: 2, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.SetChunkSize{ + Value: 65536, + }, msg) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 1, + Arguments: []interface{}{ + amf0.Object{ + {Key: "fmsVer", Value: "LNX 9,0,124,2"}, + {Key: "capabilities", Value: float64(31)}, + }, + amf0.Object{ + {Key: "level", Value: "status"}, + {Key: "code", Value: "NetConnection.Connect.Success"}, + {Key: "description", Value: "Connection succeeded."}, + {Key: "objectEncoding", Value: float64(0)}, + }, + }, + }, msg) + + err = mrw.Write(&message.SetChunkSize{ + Value: 65536, + }) + require.NoError(t, err) + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 3, + Name: "createStream", + CommandID: 2, + Arguments: []interface{}{ + nil, + }, + }) + require.NoError(t, err) + + msg, err = mrw.Read() + require.NoError(t, err) + require.Equal(t, &message.CommandAMF0{ + ChunkStreamID: 3, + Name: "_result", + CommandID: 2, + Arguments: []interface{}{ + nil, + float64(1), + }, + }, msg) + + err = mrw.Write(&message.UserControlSetBufferLength{ + BufferLength: 0x64, + }) + require.NoError(t, err) + + var streamKey string + + switch ca { + case "stream key": + streamKey = "key" + + case "stream key and query": + streamKey = "key?key=val" + } + + err = mrw.Write(&message.CommandAMF0{ + ChunkStreamID: 4, + MessageStreamID: 0x1000000, + Name: "play", + CommandID: 0, + Arguments: []interface{}{ + nil, + streamKey, + }, + }) + require.NoError(t, err) + + <-done + }) + } +} diff --git a/internal/protocols/rtmp/writer.go b/internal/protocols/rtmp/writer.go index ca1dd99a..f0ce6a00 100644 --- a/internal/protocols/rtmp/writer.go +++ b/internal/protocols/rtmp/writer.go @@ -43,14 +43,14 @@ func mpeg1AudioChannels(m mpeg1audio.ChannelMode) bool { return m != mpeg1audio.ChannelModeMono } -// Writer is a wrapper around Conn that provides utilities to mux outgoing data. +// Writer provides functions to write outgoing data. type Writer struct { - Conn *Conn + Conn Conn VideoTrack format.Format AudioTrack format.Format } -// Initialize initializes a Writer. +// Initialize initializes Writer. func (w *Writer) Initialize() error { err := w.writeTracks() if err != nil { diff --git a/internal/protocols/rtmp/writer_test.go b/internal/protocols/rtmp/writer_test.go index 45bb3aca..1b8ab201 100644 --- a/internal/protocols/rtmp/writer_test.go +++ b/internal/protocols/rtmp/writer_test.go @@ -40,19 +40,17 @@ func TestWriteTracks(t *testing.T) { } var buf bytes.Buffer - c := &Conn{ - RW: &buf, - skipHandshake: true, + c := &dummyConn{ + rw: &buf, } - err := c.Initialize() - require.NoError(t, err) + c.initialize() w := &Writer{ Conn: c, VideoTrack: videoTrack, AudioTrack: audioTrack, } - err = w.Initialize() + err := w.Initialize() require.NoError(t, err) bc := bytecounter.NewReadWriter(&buf) diff --git a/internal/servers/rtmp/conn.go b/internal/servers/rtmp/conn.go index 1b536395..d7a98f20 100644 --- a/internal/servers/rtmp/conn.go +++ b/internal/servers/rtmp/conn.go @@ -5,7 +5,6 @@ import ( "errors" "fmt" "net" - "net/url" "strings" "sync" "time" @@ -23,14 +22,6 @@ import ( "github.com/bluenviron/mediamtx/internal/stream" ) -func pathNameAndQuery(inURL *url.URL) (string, url.Values, string) { - // remove leading and trailing slashes inserted by OBS and some other clients - tmp := strings.TrimRight(inURL.String(), "/") - ur, _ := url.Parse(tmp) - pathName := strings.TrimLeft(ur.Path, "/") - return pathName, ur.Query(), ur.RawQuery -} - type connState int const ( @@ -58,7 +49,7 @@ type conn struct { uuid uuid.UUID created time.Time mutex sync.RWMutex - rconn *rtmp.Conn + rconn *rtmp.ServerConn state connState pathName string query string @@ -137,7 +128,8 @@ func (c *conn) runInner() error { func (c *conn) runReader() error { c.nconn.SetReadDeadline(time.Now().Add(time.Duration(c.readTimeout))) c.nconn.SetWriteDeadline(time.Now().Add(time.Duration(c.writeTimeout))) - conn := &rtmp.Conn{ + + conn := &rtmp.ServerConn{ RW: c.nconn, } err := conn.Initialize() @@ -145,24 +137,30 @@ func (c *conn) runReader() error { return err } + err = conn.Accept() + if err != nil { + return err + } + c.mutex.Lock() c.rconn = conn c.mutex.Unlock() if !conn.Publish { - return c.runRead(conn) + return c.runRead() } - return c.runPublish(conn) + return c.runPublish() } -func (c *conn) runRead(conn *rtmp.Conn) error { - pathName, query, rawQuery := pathNameAndQuery(conn.URL) +func (c *conn) runRead() error { + pathName := strings.TrimLeft(c.rconn.URL.Path, "/") + query := c.rconn.URL.Query() path, stream, err := c.pathManager.AddReader(defs.PathAddReaderReq{ Author: c, AccessRequest: defs.PathAccessRequest{ Name: pathName, - Query: rawQuery, + Query: c.rconn.URL.RawQuery, Proto: auth.ProtocolRTMP, ID: &c.uuid, Credentials: &auth.Credentials{ @@ -187,10 +185,10 @@ func (c *conn) runRead(conn *rtmp.Conn) error { c.mutex.Lock() c.state = connStateRead c.pathName = pathName - c.query = rawQuery + c.query = c.rconn.URL.RawQuery c.mutex.Unlock() - err = rtmp.FromStream(stream, c, conn, c.nconn, time.Duration(c.writeTimeout)) + err = rtmp.FromStream(stream, c, c.rconn, c.nconn, time.Duration(c.writeTimeout)) if err != nil { return err } @@ -204,7 +202,7 @@ func (c *conn) runRead(conn *rtmp.Conn) error { Conf: path.SafeConf(), ExternalCmdEnv: path.ExternalCmdEnv(), Reader: c.APISourceDescribe(), - Query: rawQuery, + Query: c.rconn.URL.RawQuery, }) defer onUnreadHook() @@ -223,14 +221,15 @@ func (c *conn) runRead(conn *rtmp.Conn) error { } } -func (c *conn) runPublish(conn *rtmp.Conn) error { - pathName, query, rawQuery := pathNameAndQuery(conn.URL) +func (c *conn) runPublish() error { + pathName := strings.TrimLeft(c.rconn.URL.Path, "/") + query := c.rconn.URL.Query() path, err := c.pathManager.AddPublisher(defs.PathAddPublisherReq{ Author: c, AccessRequest: defs.PathAccessRequest{ Name: pathName, - Query: rawQuery, + Query: c.rconn.URL.RawQuery, Publish: true, Proto: auth.ProtocolRTMP, ID: &c.uuid, @@ -256,11 +255,11 @@ func (c *conn) runPublish(conn *rtmp.Conn) error { c.mutex.Lock() c.state = connStatePublish c.pathName = pathName - c.query = rawQuery + c.query = c.rconn.URL.RawQuery c.mutex.Unlock() r := &rtmp.Reader{ - Conn: conn, + Conn: c.rconn, } err = r.Initialize() if err != nil { diff --git a/internal/servers/rtmp/server_test.go b/internal/servers/rtmp/server_test.go index f065a919..79d16f81 100644 --- a/internal/servers/rtmp/server_test.go +++ b/internal/servers/rtmp/server_test.go @@ -1,8 +1,8 @@ package rtmp import ( + "context" "crypto/tls" - "net" "net/url" "os" "testing" @@ -116,27 +116,28 @@ func TestServerPublish(t *testing.T) { require.NoError(t, err) defer s.Close() - u, err := url.Parse("rtmp://127.0.0.1:1935/teststream?user=myuser&pass=mypass¶m=value") - require.NoError(t, err) + var rawURL string - nconn, err := func() (net.Conn, error) { - if encrypt == "plain" { - return net.Dial("tcp", u.Host) - } - return tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - }() - require.NoError(t, err) - defer nconn.Close() - - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: true, + if encrypt == "tls" { + rawURL += "rtmps://" + } else { + rawURL += "rtmp://" } - err = conn.Initialize() + + rawURL += "127.0.0.1:1935/teststream?user=myuser&pass=mypass¶m=value" + + u, err := url.Parse(rawURL) require.NoError(t, err) + conn := &rtmp.Client{ + URL: u, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + Publish: true, + } + err = conn.Initialize(context.Background()) + require.NoError(t, err) + defer conn.Close() + w := &rtmp.Writer{ Conn: conn, VideoTrack: test.FormatH264, @@ -247,17 +248,27 @@ func TestServerRead(t *testing.T) { require.NoError(t, err) defer s.Close() - u, err := url.Parse("rtmp://127.0.0.1:1935/teststream?user=myuser&pass=mypass¶m=value") + var rawURL string + + if encrypt == "tls" { + rawURL += "rtmps://" + } else { + rawURL += "rtmp://" + } + + rawURL += "127.0.0.1:1935/teststream?user=myuser&pass=mypass¶m=value" + + u, err := url.Parse(rawURL) require.NoError(t, err) - nconn, err := func() (net.Conn, error) { - if encrypt == "plain" { - return net.Dial("tcp", u.Host) - } - return tls.Dial("tcp", u.Host, &tls.Config{InsecureSkipVerify: true}) - }() + conn := &rtmp.Client{ + URL: u, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + Publish: false, + } + err = conn.Initialize(context.Background()) require.NoError(t, err) - defer nconn.Close() + defer conn.Close() go func() { strm.WaitRunningReader() @@ -292,15 +303,6 @@ func TestServerRead(t *testing.T) { }) }() - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: false, - } - err = conn.Initialize() - require.NoError(t, err) - r := &rtmp.Reader{ Conn: conn, } diff --git a/internal/staticsources/rtmp/source.go b/internal/staticsources/rtmp/source.go index 36e62a28..35d96df8 100644 --- a/internal/staticsources/rtmp/source.go +++ b/internal/staticsources/rtmp/source.go @@ -3,7 +3,6 @@ package rtmp import ( "context" - ctls "crypto/tls" "fmt" "net" "net/url" @@ -50,53 +49,38 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { } } - nconn, err := func() (net.Conn, error) { - ctx2, cancel2 := context.WithTimeout(params.Context, time.Duration(s.ReadTimeout)) - defer cancel2() - - if u.Scheme == "rtmp" { - return (&net.Dialer{}).DialContext(ctx2, "tcp", u.Host) - } - - return (&ctls.Dialer{ - Config: tls.ConfigForFingerprint(params.Conf.SourceFingerprint), - }).DialContext(ctx2, "tcp", u.Host) - }() - if err != nil { - return err - } + ctx, ctxCancel := context.WithCancel(context.Background()) readDone := make(chan error) go func() { - readDone <- s.runReader(u, nconn) + readDone <- s.runReader(ctx, u, params.Conf.SourceFingerprint) }() for { select { case err := <-readDone: - nconn.Close() + ctxCancel() return err case <-params.ReloadConf: case <-params.Context.Done(): - nconn.Close() + ctxCancel() <-readDone return nil } } } -func (s *Source) runReader(u *url.URL, nconn net.Conn) error { - nconn.SetReadDeadline(time.Now().Add(time.Duration(s.ReadTimeout))) - nconn.SetWriteDeadline(time.Now().Add(time.Duration(s.WriteTimeout))) - conn := &rtmp.Conn{ - RW: nconn, - Client: true, - URL: u, - Publish: false, +func (s *Source) runReader(ctx context.Context, u *url.URL, fingerprint string) error { + connectCtx, connectCtxCancel := context.WithTimeout(ctx, time.Duration(s.ReadTimeout)) + conn := &rtmp.Client{ + URL: u, + TLSConfig: tls.ConfigForFingerprint(fingerprint), + Publish: false, } - err := conn.Initialize() + err := conn.Initialize(connectCtx) + connectCtxCancel() if err != nil { return err } @@ -106,6 +90,7 @@ func (s *Source) runReader(u *url.URL, nconn net.Conn) error { } err = r.Initialize() if err != nil { + conn.Close() return err } @@ -113,10 +98,12 @@ func (s *Source) runReader(u *url.URL, nconn net.Conn) error { medias, err := rtmp.ToStream(r, &stream) if err != nil { + conn.Close() return err } if len(medias) == 0 { + conn.Close() return fmt.Errorf("no supported tracks found") } @@ -125,6 +112,7 @@ func (s *Source) runReader(u *url.URL, nconn net.Conn) error { GenerateRTPPackets: true, }) if res.Err != nil { + conn.Close() return res.Err } @@ -132,15 +120,26 @@ func (s *Source) runReader(u *url.URL, nconn net.Conn) error { stream = res.Stream - // disable write deadline to allow outgoing acknowledges - nconn.SetWriteDeadline(time.Time{}) - - for { - nconn.SetReadDeadline(time.Now().Add(time.Duration(s.ReadTimeout))) - err := r.Read() - if err != nil { - return err + readerErr := make(chan error) + go func() { + for { + conn.NetConn().SetReadDeadline(time.Now().Add(time.Duration(s.ReadTimeout))) + err := r.Read() + if err != nil { + readerErr <- err + return + } } + }() + + select { + case <-ctx.Done(): + conn.Close() + <-readerErr + return fmt.Errorf("terminated") + + case err := <-readerErr: + return err } } diff --git a/internal/staticsources/rtmp/source_test.go b/internal/staticsources/rtmp/source_test.go index 9dbb4fa2..af152c52 100644 --- a/internal/staticsources/rtmp/source_test.go +++ b/internal/staticsources/rtmp/source_test.go @@ -16,63 +16,95 @@ import ( ) func TestSource(t *testing.T) { - for _, ca := range []string{ + for _, encryption := range []string{ "plain", "tls", } { - t.Run(ca, func(t *testing.T) { - ln, err := func() (net.Listener, error) { - if ca == "plain" { - return net.Listen("tcp", "127.0.0.1:1935") + for _, auth := range []string{ + "no auth", + "auth", + } { + t.Run(encryption+"_"+auth, func(t *testing.T) { + var ln net.Listener + + if encryption == "plain" { + var err error + ln, err = net.Listen("tcp", "127.0.0.1:1935") + require.NoError(t, err) + } else { + serverCertFpath, err := test.CreateTempFile(test.TLSCertPub) + require.NoError(t, err) + defer os.Remove(serverCertFpath) + + serverKeyFpath, err := test.CreateTempFile(test.TLSCertKey) + require.NoError(t, err) + defer os.Remove(serverKeyFpath) + + var cert tls.Certificate + cert, err = tls.LoadX509KeyPair(serverCertFpath, serverKeyFpath) + require.NoError(t, err) + + ln, err = tls.Listen("tcp", "127.0.0.1:1936", &tls.Config{Certificates: []tls.Certificate{cert}}) + require.NoError(t, err) } - serverCertFpath, err := test.CreateTempFile(test.TLSCertPub) - require.NoError(t, err) - defer os.Remove(serverCertFpath) + defer ln.Close() - serverKeyFpath, err := test.CreateTempFile(test.TLSCertKey) - require.NoError(t, err) - defer os.Remove(serverKeyFpath) + go func() { + for { + nconn, err := ln.Accept() + require.NoError(t, err) + defer nconn.Close() - var cert tls.Certificate - cert, err = tls.LoadX509KeyPair(serverCertFpath, serverKeyFpath) - require.NoError(t, err) + conn := &rtmp.ServerConn{ + RW: nconn, + } + err = conn.Initialize() + require.NoError(t, err) - return tls.Listen("tcp", "127.0.0.1:1936", &tls.Config{Certificates: []tls.Certificate{cert}}) - }() - require.NoError(t, err) - defer ln.Close() + if auth == "auth" { + err = conn.CheckCredentials("myuser", "mypass") + if err != nil { + continue + } + } - go func() { - nconn, err := ln.Accept() - require.NoError(t, err) - defer nconn.Close() + err = conn.Accept() + require.NoError(t, err) - conn := &rtmp.Conn{ - RW: nconn, + w := &rtmp.Writer{ + Conn: conn, + VideoTrack: test.FormatH264, + AudioTrack: test.FormatMPEG4Audio, + } + err = w.Initialize() + require.NoError(t, err) + + err = w.WriteH264(2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}}) + require.NoError(t, err) + + err = w.WriteH264(3*time.Second, 3*time.Second, [][]byte{{5, 2, 3, 4}}) + require.NoError(t, err) + + break + } + }() + + var source string + + if encryption == "plain" { + source = "rtmp://" + } else { + source = "rtmps://" } - err = conn.Initialize() - require.NoError(t, err) - w := &rtmp.Writer{ - Conn: conn, - VideoTrack: test.FormatH264, - AudioTrack: test.FormatMPEG4Audio, + if auth == "auth" { + source += "myuser:mypass@" } - err = w.Initialize() - require.NoError(t, err) - err = w.WriteH264(2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}}) - require.NoError(t, err) + source += "localhost/teststream" - err = w.WriteH264(3*time.Second, 3*time.Second, [][]byte{{5, 2, 3, 4}}) - require.NoError(t, err) - }() - - var te *test.SourceTester - - if ca == "plain" { - te = test.NewSourceTester( + te := test.NewSourceTester( func(p defs.StaticSourceParent) defs.StaticSource { return &Source{ ReadTimeout: conf.Duration(10 * time.Second), @@ -80,28 +112,16 @@ func TestSource(t *testing.T) { Parent: p, } }, - "rtmp://localhost/teststream", - &conf.Path{}, - ) - } else { - te = test.NewSourceTester( - func(p defs.StaticSourceParent) defs.StaticSource { - return &Source{ - ReadTimeout: conf.Duration(10 * time.Second), - WriteTimeout: conf.Duration(10 * time.Second), - Parent: p, - } - }, - "rtmps://localhost/teststream", + source, &conf.Path{ SourceFingerprint: "33949E05FFFB5FF3E8AA16F8213A6251B4D9363804BA53233C4DA9A46D6F2739", }, ) - } - defer te.Close() + defer te.Close() - <-te.Unit - }) + <-te.Unit + }) + } } }