Files
Alessandro RosandGitHub b80737c122 add destFingerprint parameter (#6106)
this allows to validate self-signed certificates of forward
destinations.
2026-08-18 09:35:37 +02:00

557 lines
14 KiB
Go

package core
import (
"bytes"
"context"
"fmt"
"io"
"net"
"net/http"
"net/url"
"sync/atomic"
"testing"
"time"
"github.com/bluenviron/gortmplib"
rtmpcodecs "github.com/bluenviron/gortmplib/pkg/codecs"
"github.com/google/uuid"
"github.com/pion/rtp"
pwebrtc "github.com/pion/webrtc/v4"
"github.com/stretchr/testify/require"
"github.com/bluenviron/mediamtx/internal/defs"
mtxwebrtc "github.com/bluenviron/mediamtx/internal/protocols/webrtc"
"github.com/bluenviron/mediamtx/internal/test"
)
func startRTMPForwardServer(t *testing.T) (string, <-chan [][]byte, <-chan error) {
ready := &atomic.Bool{}
ready.Store(true)
u, received, _, serverErr := startRTMPForwardServerControlled(t, ready)
return u, received, serverErr
}
func startRTMPForwardServerControlled(
t *testing.T,
ready *atomic.Bool,
) (string, <-chan [][]byte, <-chan struct{}, <-chan error) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
done := make(chan struct{})
t.Cleanup(func() {
close(done)
ln.Close()
})
received := make(chan [][]byte, 16)
connOpened := make(chan struct{}, 16)
serverErr := make(chan error, 16)
go func() {
for {
nconn, acceptErr := ln.Accept()
if acceptErr != nil {
select {
case <-done:
default:
serverErr <- acceptErr
}
return
}
if !ready.Load() {
nconn.Close()
continue
}
select {
case connOpened <- struct{}{}:
default:
}
go handleRTMPForwardConn(nconn, received, serverErr)
}
}()
return "rtmp://" + ln.Addr().String() + "/dest", received, connOpened, serverErr
}
func handleRTMPForwardConn(nconn net.Conn, received chan<- [][]byte, serverErr chan<- error) {
defer nconn.Close()
deadlineErr := nconn.SetDeadline(time.Now().Add(10 * time.Second))
if deadlineErr != nil {
serverErr <- deadlineErr
return
}
conn := &gortmplib.ServerConn{RW: nconn}
initErr := conn.Initialize()
if initErr != nil {
serverErr <- initErr
return
}
acceptConnErr := conn.AcceptConn()
if acceptConnErr != nil {
serverErr <- acceptConnErr
return
}
if !conn.Publish {
serverErr <- fmt.Errorf("connection is not publishing")
return
}
if conn.URL.Path != "/dest" {
serverErr <- fmt.Errorf("unexpected path: %s", conn.URL.Path)
return
}
r := &gortmplib.Reader{Conn: conn}
err := r.Initialize()
if err != nil {
serverErr <- err
return
}
tracks := r.Tracks()
if len(tracks) != 1 {
serverErr <- fmt.Errorf("unexpected track count: %d", len(tracks))
return
}
if _, ok := tracks[0].Codec.(*rtmpcodecs.H264); !ok {
serverErr <- fmt.Errorf("unexpected codec: %T", tracks[0].Codec)
return
}
r.OnDataH264(tracks[0], func(_ time.Duration, _ time.Duration, au [][]byte) {
for _, nalu := range au {
if bytes.Equal(nalu, []byte{5, 2, 3, 4}) {
select {
case received <- au:
default:
}
}
}
})
for {
err = r.Read()
if err != nil {
return
}
}
}
func startRTMPPublisher(
t *testing.T,
path string,
) (*gortmplib.Client, *gortmplib.Writer, *gortmplib.Track) {
u, err := url.Parse("rtmp://127.0.0.1:1935/" + path)
require.NoError(t, err)
source := &gortmplib.Client{
URL: u,
Publish: true,
}
err = source.Initialize(context.Background())
require.NoError(t, err)
track := &gortmplib.Track{
Codec: &rtmpcodecs.H264{
SPS: test.FormatH264.SPS,
PPS: test.FormatH264.PPS,
},
}
w := &gortmplib.Writer{
Conn: source,
Tracks: []*gortmplib.Track{track},
}
err = w.Initialize()
require.NoError(t, err)
return source, w, track
}
func waitRTMPForwardFrame(
t *testing.T,
w *gortmplib.Writer,
track *gortmplib.Track,
received <-chan [][]byte,
serverErr <-chan error,
) {
t.Helper()
ticker := time.NewTicker(100 * time.Millisecond)
defer ticker.Stop()
timer := time.NewTimer(10 * time.Second)
defer timer.Stop()
for {
select {
case au := <-received:
require.Contains(t, au, []byte{5, 2, 3, 4})
return
case err := <-serverErr:
require.NoError(t, err)
case <-ticker.C:
err := w.WriteH264(track, 2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}})
require.NoError(t, err)
case <-timer.C:
t.Fatal("timed out waiting for RTMP forwarded frame")
}
}
}
func startWHIPForwardServer(
t *testing.T,
expectedBearerToken string,
) (string, <-chan struct{}, <-chan error) {
pc := &mtxwebrtc.PeerConnection{
LocalRandomUDP: true,
IPsFromInterfaces: true,
Log: test.NilLogger,
}
err := pc.Start()
require.NoError(t, err)
t.Cleanup(func() {
pc.Close()
})
received := make(chan struct{}, 16)
serverErr := make(chan error, 16)
httpServ := &http.Server{
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if expectedBearerToken != "" {
require.Equal(t, "Bearer "+expectedBearerToken, r.Header.Get("Authorization"))
}
switch {
case r.Method == http.MethodOptions && r.URL.Path == "/teststream/whip":
w.Header().Set("Access-Control-Allow-Methods", "OPTIONS, GET, POST, PATCH, DELETE")
w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, If-Match")
w.WriteHeader(http.StatusNoContent)
case r.Method == http.MethodPost && r.URL.Path == "/teststream/whip":
require.Equal(t, "application/sdp", r.Header.Get("Content-Type"))
body, err2 := io.ReadAll(r.Body)
require.NoError(t, err2)
offer := &pwebrtc.SessionDescription{
Type: pwebrtc.SDPTypeOffer,
SDP: string(body),
}
answer, err2 := pc.CreateFullAnswer(offer, false)
require.NoError(t, err2)
w.Header().Set("Content-Type", "application/sdp")
w.Header().Set("ETag", "test_etag")
w.Header().Set("Location", "/teststream/whip/sessionid")
w.WriteHeader(http.StatusCreated)
_, err2 = w.Write([]byte(answer.SDP))
require.NoError(t, err2)
go func() {
err3 := pc.WaitUntilConnected(10 * time.Second)
if err3 != nil {
serverErr <- err3
return
}
err3 = pc.GatherInboundTracks(2 * time.Second)
if err3 != nil {
serverErr <- err3
return
}
if len(pc.InboundTracks()) != 1 {
serverErr <- fmt.Errorf("unexpected track count: %d", len(pc.InboundTracks()))
return
}
pc.InboundTracks()[0].OnPacketRTP = func(_ *rtp.Packet) {
select {
case received <- struct{}{}:
default:
}
}
pc.StartReading()
}()
case r.URL.Path == "/teststream/whip/sessionid" && r.Method == http.MethodPatch:
w.WriteHeader(http.StatusNoContent)
case r.URL.Path == "/teststream/whip/sessionid" && r.Method == http.MethodDelete:
w.WriteHeader(http.StatusOK)
default:
serverErr <- fmt.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusBadRequest)
}
}),
}
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
go httpServ.Serve(ln)
t.Cleanup(func() {
httpServ.Shutdown(context.Background())
})
return "whip://" + ln.Addr().String() + "/teststream/whip", received, serverErr
}
func waitWHIPForwardFrame(
t *testing.T,
w *gortmplib.Writer,
track *gortmplib.Track,
received <-chan struct{},
serverErr <-chan error,
) {
t.Helper()
ticker := time.NewTicker(100 * time.Millisecond)
defer ticker.Stop()
timer := time.NewTimer(10 * time.Second)
defer timer.Stop()
for {
select {
case <-received:
return
case err := <-serverErr:
require.NoError(t, err)
case <-ticker.C:
err := w.WriteH264(track, 2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}})
require.NoError(t, err)
case <-timer.C:
t.Fatal("timed out waiting for WHIP forwarded frame")
}
}
}
func TestPathForwardRTMP(t *testing.T) {
dest, received, serverErr := startRTMPForwardServer(t)
p, ok := newInstance(t, "api: yes\n"+
"moq: no\n"+
"paths:\n"+
" source:\n"+
" forward:\n"+
" - dest: "+dest+"\n")
require.Equal(t, true, ok)
defer p.Close()
source, w, track := startRTMPPublisher(t, "source")
defer source.Close()
tr := &http.Transport{}
defer tr.CloseIdleConnections()
hc := &http.Client{Transport: tr}
err := w.WriteH264(track, 2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}})
require.NoError(t, err)
require.Eventually(t, func() bool {
var path struct {
Ready bool `json:"ready"`
}
httpRequest(t, hc, http.MethodGet, "http://localhost:9997/v3/paths/get/source", nil, &path)
return path.Ready
}, 5*time.Second, 100*time.Millisecond)
var list defs.APIForwardDestList
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/list?path=source", nil, &list)
require.Len(t, list.Items, 1)
added := list.Items[0]
require.Equal(t, dest, added.Conf.Dest)
require.Equal(t, defs.APIForwardDestProtocolRTMP, added.Protocol)
require.Equal(t, 1, added.Pos)
waitRTMPForwardFrame(t, w, track, received, serverErr)
require.Eventually(t, func() bool {
var item defs.APIForwardDest
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/get?path=source&id="+added.ID.String(), nil, &item)
return item.State == defs.APIForwardDestStateForwarding &&
item.Protocol == defs.APIForwardDestProtocolRTMP &&
item.OutboundBytes > 0
}, 5*time.Second, 100*time.Millisecond)
}
func TestPathForwardRTMPReconnectsAfterSourceUnavailable(t *testing.T) {
dest, received, serverErr := startRTMPForwardServer(t)
p, ok := newInstance(t, "api: yes\n"+
"moq: no\n"+
"paths:\n"+
" source:\n"+
" forward:\n"+
" - dest: "+dest+"\n")
require.Equal(t, true, ok)
defer p.Close()
tr := &http.Transport{}
defer tr.CloseIdleConnections()
hc := &http.Client{Transport: tr}
var id uuid.UUID
require.Eventually(t, func() bool {
var list defs.APIForwardDestList
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/list?path=source", nil, &list)
if list.ItemCount != 1 || list.Items[0].State != defs.APIForwardDestStateIdle {
return false
}
id = list.Items[0].ID
return true
}, 7*time.Second, 100*time.Millisecond)
source, w, track := startRTMPPublisher(t, "source")
waitRTMPForwardFrame(t, w, track, received, serverErr)
source.Close()
require.Eventually(t, func() bool {
var list defs.APIForwardDestList
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/list?path=source", nil, &list)
return list.ItemCount == 1 && list.Items[0].ID == id
}, 5*time.Second, 100*time.Millisecond)
for {
select {
case <-received:
default:
goto drained
}
}
drained:
source, w, track = startRTMPPublisher(t, "source")
defer source.Close()
waitRTMPForwardFrame(t, w, track, received, serverErr)
var item defs.APIForwardDest
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/get?path=source&id="+id.String(), nil, &item)
require.Equal(t, id, item.ID)
require.Equal(t, defs.APIForwardDestStateForwarding, item.State)
require.Greater(t, item.OutboundBytes, uint64(0))
}
func TestPathForwardRTMPReconnectsAfterDestinationUnavailable(t *testing.T) {
ready := &atomic.Bool{}
dest, received, _, serverErr := startRTMPForwardServerControlled(t, ready)
p, ok := newInstance(t, "api: yes\n"+
"moq: no\n"+
"paths:\n"+
" source:\n"+
" forward:\n"+
" - dest: "+dest+"\n")
require.Equal(t, true, ok)
defer p.Close()
source, w, track := startRTMPPublisher(t, "source")
defer source.Close()
tr := &http.Transport{}
defer tr.CloseIdleConnections()
hc := &http.Client{Transport: tr}
var id uuid.UUID
require.Eventually(t, func() bool {
var list defs.APIForwardDestList
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/list?path=source", nil, &list)
if list.ItemCount != 1 || list.Items[0].State != defs.APIForwardDestStateIdle {
return false
}
id = list.Items[0].ID
return true
}, 7*time.Second, 100*time.Millisecond)
ready.Store(true)
waitRTMPForwardFrame(t, w, track, received, serverErr)
var item defs.APIForwardDest
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/get?path=source&id="+id.String(), nil, &item)
require.Equal(t, id, item.ID)
require.Equal(t, defs.APIForwardDestStateForwarding, item.State)
require.Equal(t, defs.APIForwardDestProtocolRTMP, item.Protocol)
require.Greater(t, item.OutboundBytes, uint64(0))
}
func TestPathForwardWHIP(t *testing.T) {
const bearerToken = "mytoken"
dest, received, serverErr := startWHIPForwardServer(t, bearerToken)
p, ok := newInstance(t, "api: yes\n"+
"moq: no\n"+
"paths:\n"+
" source:\n"+
" forward:\n"+
" - dest: "+dest+"\n"+
" whipBearerToken: "+bearerToken+"\n")
require.Equal(t, true, ok)
defer p.Close()
source, w, track := startRTMPPublisher(t, "source")
defer source.Close()
tr := &http.Transport{}
defer tr.CloseIdleConnections()
hc := &http.Client{Transport: tr}
err := w.WriteH264(track, 2*time.Second, 2*time.Second, [][]byte{{5, 2, 3, 4}})
require.NoError(t, err)
require.Eventually(t, func() bool {
var path struct {
Ready bool `json:"ready"`
}
httpRequest(t, hc, http.MethodGet, "http://localhost:9997/v3/paths/get/source", nil, &path)
return path.Ready
}, 5*time.Second, 100*time.Millisecond)
var list defs.APIForwardDestList
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/list?path=source", nil, &list)
require.Len(t, list.Items, 1)
added := list.Items[0]
require.Equal(t, dest, added.Conf.Dest)
require.Equal(t, bearerToken, added.Conf.WHIPBearerToken)
require.Equal(t, defs.APIForwardDestProtocolWHIP, added.Protocol)
require.Equal(t, 1, added.Pos)
waitWHIPForwardFrame(t, w, track, received, serverErr)
require.Eventually(t, func() bool {
var item defs.APIForwardDest
httpRequest(t, hc, http.MethodGet,
"http://localhost:9997/v3/paths/forward/get?path=source&id="+added.ID.String(), nil, &item)
return item.State == defs.APIForwardDestStateForwarding &&
item.Protocol == defs.APIForwardDestProtocolWHIP &&
item.Conf.WHIPBearerToken == bearerToken &&
item.OutboundBytes > 0
}, 5*time.Second, 100*time.Millisecond)
}