rtmp: support connecting to sources that require standard credentials (#4530)

This commit is contained in:
Alessandro Ros
2025-05-15 14:23:03 +02:00
committed by GitHub
parent b48a0098d3
commit 1b9dfbd367
19 changed files with 2408 additions and 1404 deletions
+21 -26
View File
@@ -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,
+11 -19
View File
@@ -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,
+18 -36
View File
@@ -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++ {
+492
View File
@@ -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)
}
+421
View File
@@ -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
})
}
}
+14 -603
View File
@@ -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)
}
-542
View File
@@ -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
}
}
+1 -1
View File
@@ -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 {
+5 -5
View File
@@ -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)
+3 -3
View File
@@ -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()
+4 -6
View File
@@ -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,
+509
View File
@@ -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)
}
+728
View File
@@ -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
})
}
}
+3 -3
View File
@@ -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 {
+4 -6
View File
@@ -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)
+23 -24
View File
@@ -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 {
+37 -35
View File
@@ -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&param=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&param=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&param=value")
var rawURL string
if encrypt == "tls" {
rawURL += "rtmps://"
} else {
rawURL += "rtmp://"
}
rawURL += "127.0.0.1:1935/teststream?user=myuser&pass=mypass&param=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,
}
+35 -36
View File
@@ -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
}
}
+79 -59
View File
@@ -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
})
}
}
}