From b2dc62e13cedf9dfc0b94f515df990b952a31866 Mon Sep 17 00:00:00 2001 From: Alex McKenzie Date: Sat, 6 Jun 2026 05:37:53 +1000 Subject: [PATCH] rtmp, rtsp: support PROXY protocol (#5754) Support PROXY protocol v1/v2 on RTMP, RTMPS, RTSP, and RTSPS TCP listeners so real client IPs are visible when running behind L4 proxies (nginx stream, HAProxy, AWS NLB). --------- Co-authored-by: aler9 <46489434+aler9@users.noreply.github.com> --- api/openapi.yaml | 8 + go.mod | 1 + go.sum | 2 + internal/conf/conf.go | 16 +- internal/core/core.go | 8 + internal/packetdumper/conn.go | 8 +- internal/packetdumper/dial_tls_context.go | 2 +- internal/packetdumper/listen.go | 23 -- internal/packetdumper/listen_packet.go | 30 -- internal/packetdumper/listener.go | 22 +- internal/packetdumper/listener_test.go | 44 +++ internal/packetdumper/packet_conn.go | 50 +-- internal/packetdumper/packet_conn_test.go | 14 +- internal/packetdumper/tls_listen.go | 24 -- internal/packetdumper/tls_listener.go | 24 +- internal/packetdumper/tls_listener_test.go | 64 ++++ internal/protocols/httpp/server.go | 42 ++- internal/protocols/proxy/listener.go | 49 +++ internal/protocols/proxy/listener_test.go | 135 ++++++++ internal/servers/rtmp/server.go | 65 ++-- internal/servers/rtmp/server_test.go | 334 +++++++++++--------- internal/servers/rtsp/server.go | 92 +++++- internal/servers/rtsp/server_test.go | 340 ++++++++++++--------- internal/staticsources/mpegts/source.go | 40 ++- internal/staticsources/rtp/source.go | 25 +- internal/staticsources/rtsp/source.go | 28 +- mediamtx.yml | 8 + 27 files changed, 1009 insertions(+), 489 deletions(-) delete mode 100644 internal/packetdumper/listen.go delete mode 100644 internal/packetdumper/listen_packet.go create mode 100644 internal/packetdumper/listener_test.go delete mode 100644 internal/packetdumper/tls_listen.go create mode 100644 internal/packetdumper/tls_listener_test.go create mode 100644 internal/protocols/proxy/listener.go create mode 100644 internal/protocols/proxy/listener_test.go diff --git a/api/openapi.yaml b/api/openapi.yaml index 029b3c75..af734f1c 100644 --- a/api/openapi.yaml +++ b/api/openapi.yaml @@ -449,6 +449,10 @@ components: type: array items: $ref: "#/components/schemas/RTSPAuthMethod" + rtspTrustedProxies: + type: array + items: + type: string rtspUDPReadBufferSize: type: integer format: uint64 @@ -472,6 +476,10 @@ components: type: string rtmpServerCert: type: string + rtmpTrustedProxies: + type: array + items: + type: string # HLS server hls: diff --git a/go.mod b/go.mod index 27a10d38..a9ee1440 100644 --- a/go.mod +++ b/go.mod @@ -37,6 +37,7 @@ require ( github.com/pion/sdp/v3 v3.0.18 github.com/pion/transport/v4 v4.0.2 github.com/pion/webrtc/v4 v4.2.14 + github.com/pires/go-proxyproto v0.12.0 github.com/quic-go/quic-go v0.59.1 github.com/quic-go/webtransport-go v0.10.0 github.com/stretchr/testify v1.11.1 diff --git a/go.sum b/go.sum index 3969e5bc..df482b55 100644 --- a/go.sum +++ b/go.sum @@ -199,6 +199,8 @@ github.com/pion/turn/v5 v5.0.8 h1:pZUCtmwWCMkrRKqh/8pL3WoGADXBe0/lOPkN7oqFjK8= github.com/pion/turn/v5 v5.0.8/go.mod h1:1VwvxElZaOdJU0liJ/WUSm/Tsh+n2OxS5ISSDxgOWxU= github.com/pion/webrtc/v4 v4.2.14 h1:Q6zMs+fSDsYuhZcNlvFGBxCOMHVV9oYcDa6O9/HIGTc= github.com/pion/webrtc/v4 v4.2.14/go.mod h1:87NVKP86+g4OMrRxWhjWfUjeXP4JrV6RTlUrIW+/Jak= +github.com/pires/go-proxyproto v0.12.0 h1:TTCxD66dU898tahivkqc3hoceZp7P44FnorWyo9d5vM= +github.com/pires/go-proxyproto v0.12.0/go.mod h1:qUvfqUMEoX7T8g0q7TQLDnhMjdTrxnG0hvpMn+7ePNI= github.com/pjbgf/sha1cd v0.6.0 h1:3WJ8Wz8gvDz29quX1OcEmkAlUg9diU4GxJHqs0/XiwU= github.com/pjbgf/sha1cd v0.6.0/go.mod h1:lhpGlyHLpQZoxMv8HcgXvZEhcGs0PG/vsZnEJ7H0iCM= github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= diff --git a/internal/conf/conf.go b/internal/conf/conf.go index 546cd19e..ac4fe73f 100644 --- a/internal/conf/conf.go +++ b/internal/conf/conf.go @@ -335,16 +335,18 @@ type Conf struct { RTSPServerCert string `json:"rtspServerCert"` AuthMethods *RTSPAuthMethods `json:"authMethods,omitempty" deprecated:"true"` RTSPAuthMethods RTSPAuthMethods `json:"rtspAuthMethods"` + RTSPTrustedProxies IPNetworks `json:"rtspTrustedProxies"` RTSPUDPReadBufferSize *uint `json:"rtspUDPReadBufferSize,omitempty" deprecated:"true"` // RTMP server - RTMP bool `json:"rtmp"` - RTMPDisable *bool `json:"rtmpDisable,omitempty" deprecated:"true"` - RTMPEncryption Encryption `json:"rtmpEncryption"` - RTMPAddress string `json:"rtmpAddress"` - RTMPSAddress string `json:"rtmpsAddress"` - RTMPServerKey string `json:"rtmpServerKey"` - RTMPServerCert string `json:"rtmpServerCert"` + RTMP bool `json:"rtmp"` + RTMPDisable *bool `json:"rtmpDisable,omitempty" deprecated:"true"` + RTMPEncryption Encryption `json:"rtmpEncryption"` + RTMPAddress string `json:"rtmpAddress"` + RTMPSAddress string `json:"rtmpsAddress"` + RTMPServerKey string `json:"rtmpServerKey"` + RTMPServerCert string `json:"rtmpServerCert"` + RTMPTrustedProxies IPNetworks `json:"rtmpTrustedProxies"` // HLS server HLS bool `json:"hls"` diff --git a/internal/core/core.go b/internal/core/core.go index 536eab1a..f9494473 100644 --- a/internal/core/core.go +++ b/internal/core/core.go @@ -491,6 +491,7 @@ func (p *Core) createResources(initial bool) error { ServerCert: "", ServerKey: "", RTSPAddress: p.conf.RTSPAddress, + TrustedProxies: p.conf.RTSPTrustedProxies, Transports: p.conf.RTSPTransports, RunOnConnect: p.conf.RunOnConnect, RunOnConnectRestart: p.conf.RunOnConnectRestart, @@ -534,6 +535,7 @@ func (p *Core) createResources(initial bool) error { ServerCert: p.conf.RTSPServerCert, ServerKey: p.conf.RTSPServerKey, RTSPAddress: p.conf.RTSPAddress, + TrustedProxies: p.conf.RTSPTrustedProxies, Transports: p.conf.RTSPTransports, RunOnConnect: p.conf.RunOnConnect, RunOnConnectRestart: p.conf.RunOnConnectRestart, @@ -563,6 +565,7 @@ func (p *Core) createResources(initial bool) error { ServerCert: "", ServerKey: "", RTSPAddress: p.conf.RTSPAddress, + TrustedProxies: p.conf.RTMPTrustedProxies, RunOnConnect: p.conf.RunOnConnect, RunOnConnectRestart: p.conf.RunOnConnectRestart, RunOnDisconnect: p.conf.RunOnDisconnect, @@ -591,6 +594,7 @@ func (p *Core) createResources(initial bool) error { ServerKey: p.conf.RTMPServerKey, DumpPackets: p.conf.DumpPackets, RTSPAddress: p.conf.RTSPAddress, + TrustedProxies: p.conf.RTMPTrustedProxies, RunOnConnect: p.conf.RunOnConnect, RunOnConnectRestart: p.conf.RunOnConnectRestart, RunOnDisconnect: p.conf.RunOnDisconnect, @@ -876,6 +880,7 @@ func (p *Core) closeResources(newConf *conf.Conf, calledByAPI bool) { newConf.MulticastRTPPort != p.conf.MulticastRTPPort || newConf.MulticastRTCPPort != p.conf.MulticastRTCPPort || !reflect.DeepEqual(newConf.RTSPTransports, p.conf.RTSPTransports) || + !reflect.DeepEqual(newConf.RTSPTrustedProxies, p.conf.RTSPTrustedProxies) || newConf.RunOnConnect != p.conf.RunOnConnect || newConf.RunOnConnectRestart != p.conf.RunOnConnectRestart || newConf.RunOnDisconnect != p.conf.RunOnDisconnect || @@ -898,6 +903,7 @@ func (p *Core) closeResources(newConf *conf.Conf, calledByAPI bool) { newConf.RTSPServerKey != p.conf.RTSPServerKey || newConf.RTSPAddress != p.conf.RTSPAddress || !reflect.DeepEqual(newConf.RTSPTransports, p.conf.RTSPTransports) || + !reflect.DeepEqual(newConf.RTSPTrustedProxies, p.conf.RTSPTrustedProxies) || newConf.RunOnConnect != p.conf.RunOnConnect || newConf.RunOnConnectRestart != p.conf.RunOnConnectRestart || newConf.RunOnDisconnect != p.conf.RunOnDisconnect || @@ -913,6 +919,7 @@ func (p *Core) closeResources(newConf *conf.Conf, calledByAPI bool) { newConf.ReadTimeout != p.conf.ReadTimeout || newConf.WriteTimeout != p.conf.WriteTimeout || newConf.RTSPAddress != p.conf.RTSPAddress || + !reflect.DeepEqual(newConf.RTMPTrustedProxies, p.conf.RTMPTrustedProxies) || newConf.RunOnConnect != p.conf.RunOnConnect || newConf.RunOnConnectRestart != p.conf.RunOnConnectRestart || newConf.RunOnDisconnect != p.conf.RunOnDisconnect || @@ -930,6 +937,7 @@ func (p *Core) closeResources(newConf *conf.Conf, calledByAPI bool) { newConf.RTMPServerCert != p.conf.RTMPServerCert || newConf.RTMPServerKey != p.conf.RTMPServerKey || newConf.RTSPAddress != p.conf.RTSPAddress || + !reflect.DeepEqual(newConf.RTMPTrustedProxies, p.conf.RTMPTrustedProxies) || newConf.RunOnConnect != p.conf.RunOnConnect || newConf.RunOnConnectRestart != p.conf.RunOnConnectRestart || newConf.RunOnDisconnect != p.conf.RunOnDisconnect || diff --git a/internal/packetdumper/conn.go b/internal/packetdumper/conn.go index 36e294a1..1e16c5ee 100644 --- a/internal/packetdumper/conn.go +++ b/internal/packetdumper/conn.go @@ -8,6 +8,7 @@ import ( "net" "os" "sync" + "sync/atomic" "time" "github.com/google/gopacket" @@ -72,7 +73,7 @@ type conn struct { Conn net.Conn ServerSide bool - expectingSecrets int + expectingSecrets atomic.Int32 f *os.File pw *pcapgo.NgWriter once sync.Once @@ -157,7 +158,7 @@ func (c *conn) run() { } func (c *conn) processEntry(e dumpEntry) { - if c.expectingSecrets > 0 && e.direction != dirSecret { + if c.expectingSecrets.Load() > 0 && e.direction != dirSecret { c.delayed = append(c.delayed, e) return } @@ -189,8 +190,7 @@ func (c *conn) processEntry(e dumpEntry) { c.pw.Flush() //nolint:errcheck writeDecryptionSecretsBlock(c.f, e.data) - c.expectingSecrets-- - if c.expectingSecrets == 0 { + if c.expectingSecrets.Add(-1) == 0 { for _, e2 := range c.delayed { c.processEntry(e2) } diff --git a/internal/packetdumper/dial_tls_context.go b/internal/packetdumper/dial_tls_context.go index dfa383fd..4c67e300 100644 --- a/internal/packetdumper/dial_tls_context.go +++ b/internal/packetdumper/dial_tls_context.go @@ -36,7 +36,7 @@ func (t *DialTLSContext) Do(ctx context.Context, network, addr string) (net.Conn } pdConn := netConn.(*conn) - pdConn.expectingSecrets = 4 + pdConn.expectingSecrets.Store(4) tlsConfig.KeyLogWriter = &connKeyLogWriter{c: pdConn} return tls.Client(netConn, tlsConfig), nil diff --git a/internal/packetdumper/listen.go b/internal/packetdumper/listen.go deleted file mode 100644 index feef1786..00000000 --- a/internal/packetdumper/listen.go +++ /dev/null @@ -1,23 +0,0 @@ -package packetdumper - -import ( - "net" -) - -// Listen is a wrapper around net.Listen that dumps packets to disk. -type Listen struct { - Prefix string -} - -// Do mimics net.Listen. -func (l *Listen) Do(network, address string) (net.Listener, error) { - netListener, err := net.Listen(network, address) - if err != nil { - return nil, err - } - - return &listener{ - Prefix: l.Prefix, - Listener: netListener, - }, nil -} diff --git a/internal/packetdumper/listen_packet.go b/internal/packetdumper/listen_packet.go deleted file mode 100644 index f90690e7..00000000 --- a/internal/packetdumper/listen_packet.go +++ /dev/null @@ -1,30 +0,0 @@ -package packetdumper - -import ( - "net" -) - -// ListenPacket is a wrapper around net.ListenPacket that dumps packets to disk. -type ListenPacket struct { - Prefix string -} - -// Do mimics net.ListenPacket. -func (l *ListenPacket) Do(network, address string) (net.PacketConn, error) { - netPacketConn, err := net.ListenPacket(network, address) - if err != nil { - return nil, err - } - - pdPacketConn := &packetConn{ - Prefix: l.Prefix, - PacketConn: netPacketConn, - } - err = pdPacketConn.Initialize() - if err != nil { - netPacketConn.Close() - return nil, err - } - - return pdPacketConn, nil -} diff --git a/internal/packetdumper/listener.go b/internal/packetdumper/listener.go index 8ef9a576..45b7f6b2 100644 --- a/internal/packetdumper/listener.go +++ b/internal/packetdumper/listener.go @@ -2,17 +2,17 @@ package packetdumper import "net" -var _ net.Listener = (*listener)(nil) +var _ net.Listener = (*Listener)(nil) -// listener is a wrapper around a net.Listener that dumps packets to disk. -type listener struct { - Prefix string - Listener net.Listener +// Listener is a wrapper around a net.Listener that dumps packets to disk. +type Listener struct { + Wrapped net.Listener + Prefix string } // Accept implements net.Listener. -func (l *listener) Accept() (net.Conn, error) { - netConn, err := l.Listener.Accept() +func (l *Listener) Accept() (net.Conn, error) { + netConn, err := l.Wrapped.Accept() if err != nil { return nil, err } @@ -32,11 +32,11 @@ func (l *listener) Accept() (net.Conn, error) { } // Close implements net.Listener. -func (l *listener) Close() error { - return l.Listener.Close() +func (l *Listener) Close() error { + return l.Wrapped.Close() } // Addr implements net.Listener. -func (l *listener) Addr() net.Addr { - return l.Listener.Addr() +func (l *Listener) Addr() net.Addr { + return l.Wrapped.Addr() } diff --git a/internal/packetdumper/listener_test.go b/internal/packetdumper/listener_test.go new file mode 100644 index 00000000..5ed6db3d --- /dev/null +++ b/internal/packetdumper/listener_test.go @@ -0,0 +1,44 @@ +package packetdumper + +import ( + "net" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestListener(t *testing.T) { + innerLn, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + prefix := filepath.Join(t.TempDir(), "capture") + ln := &Listener{ + Wrapped: innerLn, + Prefix: prefix, + } + + require.Equal(t, innerLn.Addr(), ln.Addr()) + + clientConn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + defer clientConn.Close() + + serverConn, err := ln.Accept() + require.NoError(t, err) + defer serverConn.Close() + + _, err = clientConn.Write([]byte("ping")) + require.NoError(t, err) + + buf := make([]byte, 4) + _, err = serverConn.Read(buf) + require.NoError(t, err) + require.Equal(t, []byte("ping"), buf) + + serverConn.Close() + clientConn.Close() + ln.Close() //nolint:errcheck + + checkPcapngPresence(t, prefix) +} diff --git a/internal/packetdumper/packet_conn.go b/internal/packetdumper/packet_conn.go index 1ed25b00..3e1ca28e 100644 --- a/internal/packetdumper/packet_conn.go +++ b/internal/packetdumper/packet_conn.go @@ -14,7 +14,7 @@ import ( "github.com/google/uuid" ) -var _ net.PacketConn = (*packetConn)(nil) +var _ net.PacketConn = (*PacketConn)(nil) type extendedPacketConn interface { net.PacketConn @@ -28,10 +28,10 @@ type packetDumpEntry struct { src, dst *net.UDPAddr } -// packetConn is a wrapper around net.PacketConn that dumps packets to disk. -type packetConn struct { - Prefix string - PacketConn net.PacketConn +// PacketConn is a wrapper around net.PacketConn that dumps packets to disk. +type PacketConn struct { + Prefix string + Wrapped net.PacketConn f *os.File pw *pcapgo.NgWriter @@ -43,7 +43,7 @@ type packetConn struct { } // Initialize initializes packetConn. -func (c *packetConn) Initialize() error { +func (c *PacketConn) Initialize() error { var err error c.f, err = os.Create(fmt.Sprintf("%s_%d_%s.pcapng", c.Prefix, time.Now().UnixNano(), uuid.New().String())) if err != nil { @@ -66,15 +66,15 @@ func (c *packetConn) Initialize() error { } // Close implements net.PacketConn. -func (c *packetConn) Close() error { +func (c *PacketConn) Close() error { c.once.Do(func() { close(c.terminated) }) <-c.done - return c.PacketConn.Close() + return c.Wrapped.Close() } -func (c *packetConn) run() { +func (c *PacketConn) run() { defer close(c.done) defer c.f.Close() @@ -97,7 +97,7 @@ func (c *packetConn) run() { } } -func (c *packetConn) writePacket(ntp time.Time, src, dst *net.UDPAddr, payload []byte) { +func (c *PacketConn) writePacket(ntp time.Time, src, dst *net.UDPAddr, payload []byte) { eth := &layers.Ethernet{ SrcMAC: net.HardwareAddr{0, 0, 0, 0, 0, 0}, DstMAC: net.HardwareAddr{0, 0, 0, 0, 0, 0}, @@ -130,7 +130,7 @@ func (c *packetConn) writePacket(ntp time.Time, src, dst *net.UDPAddr, payload [ }, raw) } -func (c *packetConn) enqueue(e packetDumpEntry) { +func (c *PacketConn) enqueue(e packetDumpEntry) { select { case c.queue <- e: case <-c.terminated: @@ -138,11 +138,11 @@ func (c *packetConn) enqueue(e packetDumpEntry) { } // ReadFrom implements net.PacketConn. -func (c *packetConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - n, addr, err = c.PacketConn.ReadFrom(p) +func (c *PacketConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + n, addr, err = c.Wrapped.ReadFrom(p) if n != 0 { - local := c.PacketConn.LocalAddr().(*net.UDPAddr) + local := c.Wrapped.LocalAddr().(*net.UDPAddr) remote := addr.(*net.UDPAddr) c.enqueue(packetDumpEntry{ @@ -157,11 +157,11 @@ func (c *packetConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { } // WriteTo implements net.PacketConn. -func (c *packetConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { - n, err = c.PacketConn.WriteTo(p, addr) +func (c *PacketConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { + n, err = c.Wrapped.WriteTo(p, addr) if err == nil { - local := c.PacketConn.LocalAddr().(*net.UDPAddr) + local := c.Wrapped.LocalAddr().(*net.UDPAddr) remote := addr.(*net.UDPAddr) c.enqueue(packetDumpEntry{ @@ -176,23 +176,23 @@ func (c *packetConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { } // LocalAddr implements net.PacketConn. -func (c *packetConn) LocalAddr() net.Addr { return c.PacketConn.LocalAddr() } +func (c *PacketConn) LocalAddr() net.Addr { return c.Wrapped.LocalAddr() } // SetDeadline implements net.PacketConn. -func (c *packetConn) SetDeadline(t time.Time) error { return c.PacketConn.SetDeadline(t) } +func (c *PacketConn) SetDeadline(t time.Time) error { return c.Wrapped.SetDeadline(t) } // SetReadDeadline implements net.PacketConn. -func (c *packetConn) SetReadDeadline(t time.Time) error { return c.PacketConn.SetReadDeadline(t) } +func (c *PacketConn) SetReadDeadline(t time.Time) error { return c.Wrapped.SetReadDeadline(t) } // SetWriteDeadline implements net.PacketConn. -func (c *packetConn) SetWriteDeadline(t time.Time) error { return c.PacketConn.SetWriteDeadline(t) } +func (c *PacketConn) SetWriteDeadline(t time.Time) error { return c.Wrapped.SetWriteDeadline(t) } // SetReadBuffer implements extendedPacketConn. -func (c *packetConn) SetReadBuffer(bytes int) error { - return c.PacketConn.(extendedPacketConn).SetReadBuffer(bytes) +func (c *PacketConn) SetReadBuffer(bytes int) error { + return c.Wrapped.(extendedPacketConn).SetReadBuffer(bytes) } // SyscallConn implements extendedPacketConn. -func (c *packetConn) SyscallConn() (syscall.RawConn, error) { - return c.PacketConn.(extendedPacketConn).SyscallConn() +func (c *PacketConn) SyscallConn() (syscall.RawConn, error) { + return c.Wrapped.(extendedPacketConn).SyscallConn() } diff --git a/internal/packetdumper/packet_conn_test.go b/internal/packetdumper/packet_conn_test.go index f7e059c2..e51a4636 100644 --- a/internal/packetdumper/packet_conn_test.go +++ b/internal/packetdumper/packet_conn_test.go @@ -37,7 +37,7 @@ func TestPacketConnInitialize_CreatesFile(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) c.Close() //nolint:errcheck @@ -50,7 +50,7 @@ func TestPacketConnWriteTo(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) n, err := c.WriteTo([]byte("hello world"), server.LocalAddr()) @@ -73,7 +73,7 @@ func TestPacketConnReadFrom(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) _, err := server.WriteTo([]byte("incoming data"), client.LocalAddr()) @@ -96,7 +96,7 @@ func TestPacketConnMultipleWriteRead(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) serverAddr := server.LocalAddr() @@ -140,7 +140,7 @@ func TestPacketConnCloseIdempotent(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) c.Close() //nolint:errcheck @@ -154,7 +154,7 @@ func TestPacketConnDelegatesAddrMethods(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) require.Equal(t, client.LocalAddr(), c.LocalAddr()) @@ -173,7 +173,7 @@ func TestPacketConnReadFromRecordsSource(t *testing.T) { client, server := startUDPPair(t) prefix := filepath.Join(t.TempDir(), "capture") - c := &packetConn{Prefix: prefix, PacketConn: client} + c := &PacketConn{Prefix: prefix, Wrapped: client} require.NoError(t, c.Initialize()) _, err := server.WriteTo([]byte("ping"), client.LocalAddr()) diff --git a/internal/packetdumper/tls_listen.go b/internal/packetdumper/tls_listen.go deleted file mode 100644 index 05c753a5..00000000 --- a/internal/packetdumper/tls_listen.go +++ /dev/null @@ -1,24 +0,0 @@ -package packetdumper - -import ( - "crypto/tls" - "net" -) - -// TLSListen provides a tls.Listen that also dumps TLS master secrets to disk. -type TLSListen struct { - Listen func(network string, address string) (net.Listener, error) -} - -// Do mimics tls.Listen. -func (l *TLSListen) Do(network string, laddr string, tlsConfig *tls.Config) (net.Listener, error) { - netListener, err := l.Listen(network, laddr) - if err != nil { - return nil, err - } - - return &tlsListener{ - Listener: netListener, - TLSConfig: tlsConfig, - }, nil -} diff --git a/internal/packetdumper/tls_listener.go b/internal/packetdumper/tls_listener.go index ad00815a..77488fe5 100644 --- a/internal/packetdumper/tls_listener.go +++ b/internal/packetdumper/tls_listener.go @@ -20,28 +20,34 @@ func (w *connKeyLogWriter) Write(p []byte) (int, error) { return len(p), nil } -type tlsListener struct { - Listener net.Listener +var _ net.Listener = (*TLSListener)(nil) + +// TLSListener is a wrapper around a net.Listener that dumps TLS master secrets to disk. +type TLSListener struct { + Wrapped net.Listener TLSConfig *tls.Config } -func (l *tlsListener) Close() error { - return l.Listener.Close() +// Close implements net.Listener. +func (l *TLSListener) Close() error { + return l.Wrapped.Close() } -func (l *tlsListener) Addr() net.Addr { - return l.Listener.Addr() +// Addr implements net.Listener. +func (l *TLSListener) Addr() net.Addr { + return l.Wrapped.Addr() } -func (l *tlsListener) Accept() (net.Conn, error) { - netConn, err := l.Listener.Accept() +// Accept implements net.Listener. +func (l *TLSListener) Accept() (net.Conn, error) { + netConn, err := l.Wrapped.Accept() if err != nil { return nil, err } tlsConfig := l.TLSConfig.Clone() pdConn := netConn.(*conn) - pdConn.expectingSecrets = 4 + pdConn.expectingSecrets.Store(4) tlsConfig.KeyLogWriter = &connKeyLogWriter{c: pdConn} return tls.Server(netConn, tlsConfig), nil diff --git a/internal/packetdumper/tls_listener_test.go b/internal/packetdumper/tls_listener_test.go new file mode 100644 index 00000000..b3785dfa --- /dev/null +++ b/internal/packetdumper/tls_listener_test.go @@ -0,0 +1,64 @@ +package packetdumper + +import ( + "crypto/tls" + "net" + "path/filepath" + "testing" + + "github.com/bluenviron/mediamtx/internal/test" + "github.com/stretchr/testify/require" +) + +func TestTLSListener(t *testing.T) { + cert, err := tls.X509KeyPair(test.TLSCertPub, test.TLSCertKey) + require.NoError(t, err) + + serverTLSConfig := &tls.Config{Certificates: []tls.Certificate{cert}} + + innerLn, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + prefix := filepath.Join(t.TempDir(), "capture") + + pdLn := &Listener{ + Wrapped: innerLn, + Prefix: prefix, + } + + tlsLn := &TLSListener{ + Wrapped: pdLn, + TLSConfig: serverTLSConfig, + } + + require.Equal(t, innerLn.Addr(), tlsLn.Addr()) + + clientDone := make(chan error, 1) + go func() { + clientConn, err2 := tls.Dial("tcp", tlsLn.Addr().String(), &tls.Config{InsecureSkipVerify: true}) //nolint:gosec + if err2 != nil { + clientDone <- err2 + return + } + defer clientConn.Close() //nolint:errcheck + + _, err2 = clientConn.Write([]byte("ping")) + clientDone <- err2 + }() + + serverConn, err := tlsLn.Accept() + require.NoError(t, err) + defer serverConn.Close() + + buf := make([]byte, 4) + _, err = serverConn.Read(buf) + require.NoError(t, err) + require.Equal(t, []byte("ping"), buf) + + require.NoError(t, <-clientDone) + + serverConn.Close() + tlsLn.Close() //nolint:errcheck + + checkPcapngPresence(t, prefix) +} diff --git a/internal/protocols/httpp/server.go b/internal/protocols/httpp/server.go index 80af81e7..4c5d533a 100644 --- a/internal/protocols/httpp/server.go +++ b/internal/protocols/httpp/server.go @@ -130,20 +130,38 @@ func (s *Server) Initialize() error { } } - var listen func(network string, address string) (net.Listener, error) - var tlsListen func(network string, laddr string, config *tls.Config) (net.Listener, error) + listen := func(network string, address string) (net.Listener, error) { + ln, err := net.Listen(network, address) + if err != nil { + return nil, err + } - if s.DumpPackets { - listen = (&packetdumper.Listen{ - Prefix: s.DumpPacketsPrefix, - }).Do + if s.DumpPackets { + ln = &packetdumper.Listener{ + Wrapped: ln, + Prefix: s.DumpPacketsPrefix, + } + } - tlsListen = (&packetdumper.TLSListen{ - Listen: listen, - }).Do - } else { - listen = net.Listen - tlsListen = tls.Listen + return ln, nil + } + + tlsListen := func(network string, laddr string, config *tls.Config) (net.Listener, error) { + ln, err := listen(network, laddr) + if err != nil { + return nil, err + } + + if s.DumpPackets { + ln = &packetdumper.TLSListener{ + Wrapped: ln, + TLSConfig: config, + } + } else { + ln = tls.NewListener(ln, config) + } + + return ln, nil } if tlsConfig != nil { diff --git a/internal/protocols/proxy/listener.go b/internal/protocols/proxy/listener.go new file mode 100644 index 00000000..98d6b5b8 --- /dev/null +++ b/internal/protocols/proxy/listener.go @@ -0,0 +1,49 @@ +// Package proxy provides PROXY protocol support for net.Listener. +package proxy + +import ( + "net" + + "github.com/pires/go-proxyproto" + + "github.com/bluenviron/mediamtx/internal/conf" +) + +var _ net.Listener = (*Listener)(nil) + +// Listener is a net.Listener that supports PROXY protocol. +type Listener struct { + Wrapped net.Listener + TrustedProxies conf.IPNetworks + + inner *proxyproto.Listener +} + +// Initialize initializes the listener. +func (l *Listener) Initialize() { + l.inner = &proxyproto.Listener{ + Listener: l.Wrapped, + Policy: func(upstream net.Addr) (proxyproto.Policy, error) { + tcpAddr, ok := upstream.(*net.TCPAddr) + if ok && l.TrustedProxies.Contains(tcpAddr.IP) { + return proxyproto.USE, nil + } + return proxyproto.IGNORE, nil + }, + } +} + +// Close implements net.Listener. +func (l *Listener) Close() error { + return l.Wrapped.Close() +} + +// Accept implements net.Listener. +func (l *Listener) Accept() (net.Conn, error) { + return l.inner.Accept() +} + +// Addr implements net.Listener. +func (l *Listener) Addr() net.Addr { + return l.Wrapped.Addr() +} diff --git a/internal/protocols/proxy/listener_test.go b/internal/protocols/proxy/listener_test.go new file mode 100644 index 00000000..c1dfc76a --- /dev/null +++ b/internal/protocols/proxy/listener_test.go @@ -0,0 +1,135 @@ +package proxy + +import ( + "fmt" + "net" + "testing" + + "github.com/pires/go-proxyproto" + "github.com/stretchr/testify/require" + + "github.com/bluenviron/mediamtx/internal/conf" +) + +func trustedProxies(cidrs ...string) conf.IPNetworks { + networks := make(conf.IPNetworks, 0, len(cidrs)) + for _, cidr := range cidrs { + _, ipnet, err := net.ParseCIDR(cidr) + if err != nil { + panic(err) + } + networks = append(networks, conf.IPNetwork(*ipnet)) + } + return networks +} + +func TestListenerTrustedWithHeader(t *testing.T) { + for _, version := range []byte{1, 2} { + t.Run(fmt.Sprintf("v%d", version), func(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + + wrapped := &Listener{Wrapped: ln, TrustedProxies: trustedProxies("127.0.0.1/32")} + wrapped.Initialize() + + done := make(chan struct{}) + + go func() { + defer close(done) + + conn, err2 := wrapped.Accept() + require.NoError(t, err2) + defer conn.Close() + + require.Equal(t, "192.168.1.100:1234", conn.RemoteAddr().String()) + }() + + clientConn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + defer clientConn.Close() + + header := &proxyproto.Header{ + Version: version, + Command: proxyproto.PROXY, + TransportProtocol: proxyproto.TCPv4, + SourceAddr: &net.TCPAddr{IP: net.ParseIP("192.168.1.100"), Port: 1234}, + DestinationAddr: &net.TCPAddr{IP: net.ParseIP("10.0.0.1"), Port: 1935}, + } + _, err = header.WriteTo(clientConn) + require.NoError(t, err) + + <-done + }) + } +} + +func TestListenerTrustedWithoutHeader(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + + wrapped := &Listener{Wrapped: ln, TrustedProxies: trustedProxies("127.0.0.1/32")} + wrapped.Initialize() + + done := make(chan struct{}) + + go func() { + defer close(done) + + conn, err2 := wrapped.Accept() + require.NoError(t, err2) + defer conn.Close() + + buf := make([]byte, 7) + _, err2 = conn.Read(buf) + require.NoError(t, err2) + require.Equal(t, []byte("testing"), buf) + }() + + clientConn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + defer clientConn.Close() + + _, err = clientConn.Write([]byte("testing")) + require.NoError(t, err) + + <-done +} + +func TestListenerUntrustedIgnoresHeader(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + + wrapped := &Listener{Wrapped: ln, TrustedProxies: trustedProxies("10.0.0.0/8")} + wrapped.Initialize() + + done := make(chan struct{}) + + go func() { + defer close(done) + + conn, err2 := wrapped.Accept() + require.NoError(t, err2) + defer conn.Close() + + require.Equal(t, "127.0.0.1", conn.RemoteAddr().(*net.TCPAddr).IP.String()) + }() + + clientConn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + defer clientConn.Close() + + header := &proxyproto.Header{ + Version: 1, + Command: proxyproto.PROXY, + TransportProtocol: proxyproto.TCPv4, + SourceAddr: &net.TCPAddr{IP: net.ParseIP("192.168.1.100"), Port: 1234}, + DestinationAddr: &net.TCPAddr{IP: net.ParseIP("10.0.0.1"), Port: 1935}, + } + _, err = header.WriteTo(clientConn) + require.NoError(t, err) + + <-done +} diff --git a/internal/servers/rtmp/server.go b/internal/servers/rtmp/server.go index 791ab13f..4694bf47 100644 --- a/internal/servers/rtmp/server.go +++ b/internal/servers/rtmp/server.go @@ -19,6 +19,7 @@ import ( "github.com/bluenviron/mediamtx/internal/externalcmd" "github.com/bluenviron/mediamtx/internal/logger" "github.com/bluenviron/mediamtx/internal/packetdumper" + "github.com/bluenviron/mediamtx/internal/protocols/proxy" "github.com/bluenviron/mediamtx/internal/restrictnetwork" ) @@ -81,6 +82,7 @@ type Server struct { ServerCert string ServerKey string RTSPAddress string + TrustedProxies conf.IPNetworks RunOnConnect string RunOnConnectRestart bool RunOnDisconnect string @@ -107,27 +109,54 @@ type Server struct { // Initialize initializes the server. func (s *Server) Initialize() error { - var listen func(network string, address string) (net.Listener, error) - var tlsListen func(network string, laddr string, config *tls.Config) (net.Listener, error) - - if s.DumpPackets { - var proto string - if s.Encryption { - proto = "rtmps" - } else { - proto = "rtmp" + listen := func(network, address string) (net.Listener, error) { + ln, err := net.Listen(network, address) + if err != nil { + return nil, err } - listen = (&packetdumper.Listen{ - Prefix: proto + "_server_conn", - }).Do + if s.DumpPackets { + var proto string + if s.Encryption { + proto = "rtmps" + } else { + proto = "rtmp" + } - tlsListen = (&packetdumper.TLSListen{ - Listen: listen, - }).Do - } else { - listen = net.Listen - tlsListen = tls.Listen + ln = &packetdumper.Listener{ + Wrapped: ln, + Prefix: proto + "_server_conn", + } + } + + if len(s.TrustedProxies) > 0 { + pl := &proxy.Listener{ + Wrapped: ln, + TrustedProxies: s.TrustedProxies, + } + pl.Initialize() + ln = pl + } + + return ln, nil + } + + tlsListen := func(network string, laddr string, config *tls.Config) (net.Listener, error) { + ln, err := listen(network, laddr) + if err != nil { + return nil, err + } + + if s.DumpPackets { + ln = &packetdumper.TLSListener{ + Wrapped: ln, + TLSConfig: config, + } + } else { + ln = tls.NewListener(ln, config) + } + + return ln, nil } if s.Encryption { diff --git a/internal/servers/rtmp/server_test.go b/internal/servers/rtmp/server_test.go index 5aea9c2e..2bef6122 100644 --- a/internal/servers/rtmp/server_test.go +++ b/internal/servers/rtmp/server_test.go @@ -3,6 +3,7 @@ package rtmp import ( "context" "crypto/tls" + "net" "net/url" "testing" "time" @@ -10,6 +11,8 @@ import ( "github.com/bluenviron/gortmplib" "github.com/bluenviron/gortmplib/pkg/codecs" "github.com/bluenviron/gortsplib/v5/pkg/description" + "github.com/pires/go-proxyproto" + "github.com/bluenviron/mediamtx/internal/conf" "github.com/bluenviron/mediamtx/internal/defs" "github.com/bluenviron/mediamtx/internal/externalcmd" @@ -44,168 +47,201 @@ func TestServerPublish(t *testing.T) { "plain", "tls", } { - t.Run(encrypt, func(t *testing.T) { - var serverCertFpath string - var serverKeyFpath string + for _, proxy := range []string{ + "no_proxy", + "proxy", + } { + t.Run(encrypt+"_"+proxy, func(t *testing.T) { + var serverCertFpath string + var serverKeyFpath string - if encrypt == "tls" { - serverCertFpath = test.CreateTempFile(t, test.TLSCertPub) - serverKeyFpath = test.CreateTempFile(t, test.TLSCertKey) - } + if encrypt == "tls" { + serverCertFpath = test.CreateTempFile(t, test.TLSCertPub) + serverKeyFpath = test.CreateTempFile(t, test.TLSCertKey) + } - var strm *stream.Stream - var reader *stream.Reader - defer func() { - strm.RemoveReader(reader) - }() - dataReceived := make(chan struct{}) - n := 0 + _, ipnet, err := net.ParseCIDR("127.0.0.1/32") + require.NoError(t, err) + trustedProxies := conf.IPNetworks{conf.IPNetwork(*ipnet)} - pathManager := &test.PathManager{ - AddPublisherImpl: func(req defs.PathAddPublisherReq) (*defs.PathAddPublisherRes, error) { - require.Equal(t, "teststream", req.AccessRequest.Name) - require.Equal(t, "user=myuser&pass=mypass¶m=value", req.AccessRequest.Query) - require.Equal(t, "myuser", req.AccessRequest.Credentials.User) - require.Equal(t, "mypass", req.AccessRequest.Credentials.Pass) + var strm *stream.Stream + var reader *stream.Reader + defer func() { + strm.RemoveReader(reader) + }() + dataReceived := make(chan struct{}) + n := 0 - strm = &stream.Stream{ - Desc: req.Desc, - WriteQueueSize: 512, - RTPMaxPayloadSize: 1450, - Parent: test.NilLogger, - } - err := strm.Initialize() - require.NoError(t, err) + pathManager := &test.PathManager{ + AddPublisherImpl: func(req defs.PathAddPublisherReq) (*defs.PathAddPublisherRes, error) { + require.Equal(t, "teststream", req.AccessRequest.Name) + require.Equal(t, "user=myuser&pass=mypass¶m=value", req.AccessRequest.Query) + require.Equal(t, "myuser", req.AccessRequest.Credentials.User) + require.Equal(t, "mypass", req.AccessRequest.Credentials.Pass) - subStream := &stream.SubStream{ - Stream: strm, - UseRTPPackets: false, - } - err = subStream.Initialize() - require.NoError(t, err) + strm = &stream.Stream{ + Desc: req.Desc, + WriteQueueSize: 512, + RTPMaxPayloadSize: 1450, + Parent: test.NilLogger, + } + err2 := strm.Initialize() + require.NoError(t, err2) - reader = &stream.Reader{Parent: test.NilLogger} + subStream := &stream.SubStream{ + Stream: strm, + UseRTPPackets: false, + } + err2 = subStream.Initialize() + require.NoError(t, err2) - reader.OnData( - strm.Desc.Medias[0], - strm.Desc.Medias[0].Formats[0], - func(u *unit.Unit) error { - switch n { - case 0: - require.Equal(t, unit.PayloadH264(nil), u.Payload) + reader = &stream.Reader{Parent: test.NilLogger} - case 1: - require.Equal(t, unit.PayloadH264{ - test.FormatH264.SPS, - test.FormatH264.PPS, - {5, 2, 3, 4}, - }, u.Payload) - close(dataReceived) + reader.OnData( + strm.Desc.Medias[0], + strm.Desc.Medias[0].Formats[0], + func(u *unit.Unit) error { + switch n { + case 0: + require.Equal(t, unit.PayloadH264(nil), u.Payload) - default: - t.Errorf("should not happen") - } - n++ - return nil - }) + case 1: + require.Equal(t, unit.PayloadH264{ + test.FormatH264.SPS, + test.FormatH264.PPS, + {5, 2, 3, 4}, + }, u.Payload) + close(dataReceived) - strm.AddReader(reader) + default: + t.Errorf("should not happen") + } + n++ + return nil + }) - return &defs.PathAddPublisherRes{ - Path: &dummyPath{}, - User: req.AccessRequest.Credentials.User, - SubStream: subStream, - }, nil - }, - } + strm.AddReader(reader) - s := &Server{ - Address: "127.0.0.1:1939", - ReadTimeout: conf.Duration(10 * time.Second), - WriteTimeout: conf.Duration(10 * time.Second), - Encryption: encrypt == "tls", - ServerCert: serverCertFpath, - ServerKey: serverKeyFpath, - RTSPAddress: "", - RunOnConnect: "", - RunOnConnectRestart: false, - RunOnDisconnect: "", - ExternalCmdPool: nil, - PathManager: pathManager, - Parent: test.NilLogger, - } - err := s.Initialize() - require.NoError(t, err) - defer s.Close() - - var rawURL string - - if encrypt == "tls" { - rawURL += "rtmps://" - } else { - rawURL += "rtmp://" - } - - rawURL += "127.0.0.1:1939/teststream?user=myuser&pass=mypass¶m=value" - - u, err := url.Parse(rawURL) - require.NoError(t, err) - - conn := &gortmplib.Client{ - URL: u, - TLSConfig: &tls.Config{InsecureSkipVerify: true}, - Publish: true, - } - err = conn.Initialize(context.Background()) - require.NoError(t, err) - defer conn.Close() - - w := &gortmplib.Writer{ - Conn: conn, - Tracks: []*gortmplib.Track{ - {Codec: &codecs.H264{ - SPS: test.FormatH264.SPS, - PPS: test.FormatH264.PPS, - }}, - {Codec: &codecs.MPEG4Audio{ - Config: test.FormatMPEG4Audio.Config, - }}, - }, - } - err = w.Initialize() - require.NoError(t, err) - - err = w.WriteH264( - w.Tracks[0], - 2*time.Second, 2*time.Second, [][]byte{ - {5, 2, 3, 4}, - }) - require.NoError(t, err) - - <-dataReceived - - list, err := s.APIConnsList() - require.NoError(t, err) - require.Equal(t, &defs.APIRTMPConnList{ - Items: []defs.APIRTMPConn{ - { - ID: list.Items[0].ID, - Created: list.Items[0].Created, - RemoteAddr: list.Items[0].RemoteAddr, - State: "publish", - Path: "teststream", - Query: "user=myuser&pass=mypass¶m=value", - User: "myuser", - UserAgent: list.Items[0].UserAgent, - InboundBytes: list.Items[0].InboundBytes, - OutboundBytes: list.Items[0].OutboundBytes, - OutboundFramesDiscarded: list.Items[0].OutboundFramesDiscarded, - BytesReceived: list.Items[0].BytesReceived, - BytesSent: list.Items[0].BytesSent, + return &defs.PathAddPublisherRes{ + Path: &dummyPath{}, + User: req.AccessRequest.Credentials.User, + SubStream: subStream, + }, nil }, - }, - }, list) - }) + } + + s := &Server{ + Address: "127.0.0.1:1939", + ReadTimeout: conf.Duration(10 * time.Second), + WriteTimeout: conf.Duration(10 * time.Second), + Encryption: encrypt == "tls", + ServerCert: serverCertFpath, + ServerKey: serverKeyFpath, + RTSPAddress: "", + TrustedProxies: trustedProxies, + RunOnConnect: "", + RunOnConnectRestart: false, + RunOnDisconnect: "", + ExternalCmdPool: nil, + PathManager: pathManager, + Parent: test.NilLogger, + } + err = s.Initialize() + require.NoError(t, err) + defer s.Close() + + var rawURL string + + if encrypt == "tls" { + rawURL += "rtmps://" + } else { + rawURL += "rtmp://" + } + + rawURL += "127.0.0.1:1939/teststream?user=myuser&pass=mypass¶m=value" + + u, err := url.Parse(rawURL) + require.NoError(t, err) + + var dialContext func(ctx context.Context, network, address string) (net.Conn, error) + if proxy == "proxy" { + dialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + c, err2 := (&net.Dialer{}).DialContext(ctx, network, address) + if err2 != nil { + return nil, err2 + } + header := &proxyproto.Header{ + Version: 1, + Command: proxyproto.PROXY, + TransportProtocol: proxyproto.TCPv4, + SourceAddr: &net.TCPAddr{IP: net.ParseIP("192.168.1.100"), Port: 1234}, + DestinationAddr: &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: 1939}, + } + _, err2 = header.WriteTo(c) + if err2 != nil { + return nil, err2 + } + return c, nil + } + } + + conn := &gortmplib.Client{ + URL: u, + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + Publish: true, + DialContext: dialContext, + } + err = conn.Initialize(context.Background()) + require.NoError(t, err) + defer conn.Close() + + w := &gortmplib.Writer{ + Conn: conn, + Tracks: []*gortmplib.Track{ + {Codec: &codecs.H264{ + SPS: test.FormatH264.SPS, + PPS: test.FormatH264.PPS, + }}, + {Codec: &codecs.MPEG4Audio{ + Config: test.FormatMPEG4Audio.Config, + }}, + }, + } + err = w.Initialize() + require.NoError(t, err) + + err = w.WriteH264( + w.Tracks[0], + 2*time.Second, 2*time.Second, [][]byte{ + {5, 2, 3, 4}, + }) + require.NoError(t, err) + + <-dataReceived + + list, err := s.APIConnsList() + require.NoError(t, err) + require.Equal(t, &defs.APIRTMPConnList{ + Items: []defs.APIRTMPConn{ + { + ID: list.Items[0].ID, + Created: list.Items[0].Created, + RemoteAddr: list.Items[0].RemoteAddr, + State: "publish", + Path: "teststream", + Query: "user=myuser&pass=mypass¶m=value", + User: "myuser", + UserAgent: list.Items[0].UserAgent, + InboundBytes: list.Items[0].InboundBytes, + OutboundBytes: list.Items[0].OutboundBytes, + OutboundFramesDiscarded: list.Items[0].OutboundFramesDiscarded, + BytesReceived: list.Items[0].BytesReceived, + BytesSent: list.Items[0].BytesSent, + }, + }, + }, list) + }) + } } } diff --git a/internal/servers/rtsp/server.go b/internal/servers/rtsp/server.go index f6f31fce..27220bce 100644 --- a/internal/servers/rtsp/server.go +++ b/internal/servers/rtsp/server.go @@ -6,6 +6,7 @@ import ( "crypto/tls" "errors" "fmt" + "net" "reflect" "sort" "strings" @@ -24,6 +25,7 @@ import ( "github.com/bluenviron/mediamtx/internal/externalcmd" "github.com/bluenviron/mediamtx/internal/logger" "github.com/bluenviron/mediamtx/internal/packetdumper" + "github.com/bluenviron/mediamtx/internal/protocols/proxy" ) // ErrConnNotFound is returned when a connection is not found. @@ -105,6 +107,7 @@ type Server struct { ServerCert string ServerKey string RTSPAddress string + TrustedProxies conf.IPNetworks Transports conf.RTSPTransports RunOnConnect string RunOnConnectRestart bool @@ -166,25 +169,84 @@ func (s *Server) Initialize() error { s.srv.TLSConfig = &tls.Config{GetCertificate: s.loader.GetCertificate()} } - if s.DumpPackets { - var proto string - if s.Encryption { - proto = "rtsps" - } else { - proto = "rtsp" + s.srv.Listen = func(network, address string) (net.Listener, error) { + ln, err := net.Listen(network, address) + if err != nil { + return nil, err } - s.srv.Listen = (&packetdumper.Listen{ - Prefix: proto + "_server_conn", - }).Do + if s.DumpPackets { + var proto string + if s.Encryption { + proto = "rtsps" + } else { + proto = "rtsp" + } - s.srv.ListenPacket = (&packetdumper.ListenPacket{ - Prefix: proto + "_server_packet_conn", - }).Do + ln = &packetdumper.Listener{ + Wrapped: ln, + Prefix: proto + "_server_conn", + } + } - s.srv.TLSListen = (&packetdumper.TLSListen{ - Listen: s.srv.Listen, - }).Do + if len(s.TrustedProxies) > 0 { + pl := &proxy.Listener{ + Wrapped: ln, + TrustedProxies: s.TrustedProxies, + } + pl.Initialize() + ln = pl + } + + return ln, nil + } + + s.srv.TLSListen = func(network, laddr string, config *tls.Config) (net.Listener, error) { + ln, err := s.srv.Listen(network, laddr) + if err != nil { + return nil, err + } + + if s.DumpPackets { + ln = &packetdumper.TLSListener{ + Wrapped: ln, + TLSConfig: config, + } + } else { + ln = tls.NewListener(ln, config) + } + + return ln, nil + } + + s.srv.ListenPacket = func(network, address string) (net.PacketConn, error) { + pc, err := net.ListenPacket(network, address) + if err != nil { + return nil, err + } + + if s.DumpPackets { + var proto string + if s.Encryption { + proto = "rtsps" + } else { + proto = "rtsp" + } + + pc2 := &packetdumper.PacketConn{ + Wrapped: pc, + Prefix: proto + "_server_packet_conn", + } + err = pc2.Initialize() + if err != nil { + pc.Close() //nolint:errcheck + return nil, err + } + + pc = pc2 + } + + return pc, nil } err := s.srv.Start() diff --git a/internal/servers/rtsp/server_test.go b/internal/servers/rtsp/server_test.go index 57fa453c..b174a942 100644 --- a/internal/servers/rtsp/server_test.go +++ b/internal/servers/rtsp/server_test.go @@ -3,7 +3,10 @@ package rtsp import ( "bufio" "bytes" + "context" + "crypto/tls" "fmt" + "net" "sync/atomic" "testing" "time" @@ -15,6 +18,8 @@ import ( "github.com/bluenviron/gortsplib/v5/pkg/format" mpegts "github.com/bluenviron/mediacommon/v2/pkg/formats/mpegts" tscodecs "github.com/bluenviron/mediacommon/v2/pkg/formats/mpegts/codecs" + "github.com/pires/go-proxyproto" + "github.com/bluenviron/mediamtx/internal/auth" "github.com/bluenviron/mediamtx/internal/conf" "github.com/bluenviron/mediamtx/internal/defs" @@ -49,159 +54,216 @@ func (p *dummyPath) RemoveReader(_ defs.PathRemoveReaderReq) { func TestServerPublish(t *testing.T) { for _, ca := range []string{"basic", "digest", "basic+digest"} { - t.Run(ca, func(t *testing.T) { - var strm *stream.Stream - var reader *stream.Reader - defer func() { - strm.RemoveReader(reader) - }() - dataReceived := make(chan struct{}) + for _, encrypt := range []string{"plain", "tls"} { + for _, proxy := range []string{"no_proxy", "proxy"} { + t.Run(ca+"_"+encrypt+"_"+proxy, func(t *testing.T) { + var serverCertFpath string + var serverKeyFpath string - n := 0 + if encrypt == "tls" { + serverCertFpath = test.CreateTempFile(t, test.TLSCertPub) + serverKeyFpath = test.CreateTempFile(t, test.TLSCertKey) + } - pathManager := &test.PathManager{ - FindPathConfImpl: func(req defs.PathFindPathConfReq) (*defs.PathFindPathConfRes, error) { - require.Equal(t, "teststream", req.AccessRequest.Name) - require.Equal(t, "param=value", req.AccessRequest.Query) + _, ipnet, err := net.ParseCIDR("127.0.0.1/32") + require.NoError(t, err) + trustedProxies := conf.IPNetworks{conf.IPNetwork(*ipnet)} - if ca == "basic" { - require.Nil(t, req.AccessRequest.CustomVerifyFunc) + var strm *stream.Stream + var reader *stream.Reader + defer func() { + strm.RemoveReader(reader) + }() + dataReceived := make(chan struct{}) - if req.AccessRequest.Credentials.User == "" && req.AccessRequest.Credentials.Pass == "" { - return nil, &auth.Error{AskCredentials: true, Wrapped: fmt.Errorf("auth error")} - } + n := 0 - require.Equal(t, "myuser", req.AccessRequest.Credentials.User) - require.Equal(t, "mypass", req.AccessRequest.Credentials.Pass) + pathManager := &test.PathManager{ + FindPathConfImpl: func(req defs.PathFindPathConfReq) (*defs.PathFindPathConfRes, error) { + require.Equal(t, "teststream", req.AccessRequest.Name) + require.Equal(t, "param=value", req.AccessRequest.Query) + + if ca == "basic" { + require.Nil(t, req.AccessRequest.CustomVerifyFunc) + + if req.AccessRequest.Credentials.User == "" && req.AccessRequest.Credentials.Pass == "" { + return nil, &auth.Error{AskCredentials: true, Wrapped: fmt.Errorf("auth error")} + } + + require.Equal(t, "myuser", req.AccessRequest.Credentials.User) + require.Equal(t, "mypass", req.AccessRequest.Credentials.Pass) + } else { + ok := req.AccessRequest.CustomVerifyFunc("myuser", "mypass") + if n == 0 { + require.False(t, ok) + n++ + return nil, &auth.Error{AskCredentials: true, Wrapped: fmt.Errorf("auth error")} + } + require.True(t, ok) + } + + return &defs.PathFindPathConfRes{Conf: &conf.Path{}, User: req.AccessRequest.Credentials.User}, nil + }, + AddPublisherImpl: func(req defs.PathAddPublisherReq) (*defs.PathAddPublisherRes, error) { + require.Equal(t, "teststream", req.AccessRequest.Name) + require.Equal(t, "param=value", req.AccessRequest.Query) + require.True(t, req.AccessRequest.SkipAuth) + + strm = &stream.Stream{ + Desc: req.Desc, + WriteQueueSize: 512, + RTPMaxPayloadSize: 1450, + Parent: test.NilLogger, + } + err2 := strm.Initialize() + require.NoError(t, err2) + + subStream := &stream.SubStream{ + Stream: strm, + UseRTPPackets: true, + } + err2 = subStream.Initialize() + require.NoError(t, err2) + + reader = &stream.Reader{Parent: test.NilLogger} + + reader.OnData( + strm.Desc.Medias[0], + strm.Desc.Medias[0].Formats[0], + func(u *unit.Unit) error { + require.Equal(t, unit.PayloadH264{ + test.FormatH264.SPS, + test.FormatH264.PPS, + {5, 2, 3, 4}, + }, u.Payload) + close(dataReceived) + return nil + }) + + strm.AddReader(reader) + + return &defs.PathAddPublisherRes{Path: &dummyPath{}, SubStream: subStream}, nil + }, + } + + var authMethods []rtspauth.VerifyMethod + switch ca { + case "basic": + authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodBasic} + case "digest": + authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodDigestMD5} + default: + authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodBasic, rtspauth.VerifyMethodDigestMD5} + } + + s := &Server{ + Address: "127.0.0.1:8557", + AuthMethods: authMethods, + ReadTimeout: conf.Duration(10 * time.Second), + WriteTimeout: conf.Duration(10 * time.Second), + WriteQueueSize: 512, + Transports: conf.RTSPTransports{gortsplib.ProtocolTCP: {}}, + Encryption: encrypt == "tls", + ServerCert: serverCertFpath, + ServerKey: serverKeyFpath, + TrustedProxies: trustedProxies, + PathManager: pathManager, + Parent: test.NilLogger, + } + err = s.Initialize() + require.NoError(t, err) + defer s.Close() + + var scheme string + if encrypt == "tls" { + scheme = "rtsps" } else { - ok := req.AccessRequest.CustomVerifyFunc("myuser", "mypass") - if n == 0 { - require.False(t, ok) - n++ - return nil, &auth.Error{AskCredentials: true, Wrapped: fmt.Errorf("auth error")} + scheme = "rtsp" + } + + var dialContext func(ctx context.Context, network, address string) (net.Conn, error) + if proxy == "proxy" { + dialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + c, err2 := (&net.Dialer{}).DialContext(ctx, network, address) + if err2 != nil { + return nil, err2 + } + header := &proxyproto.Header{ + Version: 1, + Command: proxyproto.PROXY, + TransportProtocol: proxyproto.TCPv4, + SourceAddr: &net.TCPAddr{IP: net.ParseIP("192.168.1.100"), Port: 1234}, + DestinationAddr: &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: 8557}, + } + _, err2 = header.WriteTo(c) + if err2 != nil { + return nil, err2 + } + return c, nil } - require.True(t, ok) } - return &defs.PathFindPathConfRes{Conf: &conf.Path{}, User: req.AccessRequest.Credentials.User}, nil - }, - AddPublisherImpl: func(req defs.PathAddPublisherReq) (*defs.PathAddPublisherRes, error) { - require.Equal(t, "teststream", req.AccessRequest.Name) - require.Equal(t, "param=value", req.AccessRequest.Query) - require.True(t, req.AccessRequest.SkipAuth) - - strm = &stream.Stream{ - Desc: req.Desc, - WriteQueueSize: 512, - RTPMaxPayloadSize: 1450, - Parent: test.NilLogger, + source := gortsplib.Client{ + TLSConfig: &tls.Config{InsecureSkipVerify: true}, + DialContext: dialContext, } - err := strm.Initialize() + + media0 := test.UniqueMediaH264() + + err = source.StartRecording( + scheme+"://myuser:mypass@127.0.0.1:8557/teststream?param=value", + &description.Session{Medias: []*description.Media{media0}}) + require.NoError(t, err) + defer source.Close() + + err = source.WritePacketRTP(media0, &rtp.Packet{ + Header: rtp.Header{ + Version: 2, + Marker: true, + PayloadType: 96, + SequenceNumber: 123, + Timestamp: 45343, + SSRC: 563423, + }, + Payload: []byte{5, 2, 3, 4}, + }) require.NoError(t, err) - subStream := &stream.SubStream{ - Stream: strm, - UseRTPPackets: true, - } - err = subStream.Initialize() + <-dataReceived + + list, err := s.APISessionsList() require.NoError(t, err) - - reader = &stream.Reader{Parent: test.NilLogger} - - reader.OnData( - strm.Desc.Medias[0], - strm.Desc.Medias[0].Formats[0], - func(u *unit.Unit) error { - require.Equal(t, unit.PayloadH264{ - test.FormatH264.SPS, - test.FormatH264.PPS, - {5, 2, 3, 4}, - }, u.Payload) - close(dataReceived) - return nil - }) - - strm.AddReader(reader) - - return &defs.PathAddPublisherRes{Path: &dummyPath{}, SubStream: subStream}, nil - }, + require.Equal(t, &defs.APIRTSPSessionList{ + Items: []defs.APIRTSPSession{ + { + ID: list.Items[0].ID, + Created: list.Items[0].Created, + RemoteAddr: list.Items[0].RemoteAddr, + State: "publish", + Path: "teststream", + Query: "param=value", + User: "myuser", + UserAgent: list.Items[0].UserAgent, + InboundBytes: list.Items[0].InboundBytes, + InboundRTPPackets: list.Items[0].InboundRTPPackets, + OutboundBytes: list.Items[0].OutboundBytes, + BytesReceived: list.Items[0].BytesReceived, + BytesSent: list.Items[0].BytesSent, + Conns: list.Items[0].Conns, + RTPPacketsReceived: list.Items[0].RTPPacketsReceived, + Transport: new("TCP"), + Profile: func() *string { + if encrypt == "tls" { + return new("SAVP") + } + return new("AVP") + }(), + }, + }, + }, list) + }) } - - var authMethods []rtspauth.VerifyMethod - switch ca { - case "basic": - authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodBasic} - case "digest": - authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodDigestMD5} - default: - authMethods = []rtspauth.VerifyMethod{rtspauth.VerifyMethodBasic, rtspauth.VerifyMethodDigestMD5} - } - - s := &Server{ - Address: "127.0.0.1:8557", - AuthMethods: authMethods, - ReadTimeout: conf.Duration(10 * time.Second), - WriteTimeout: conf.Duration(10 * time.Second), - WriteQueueSize: 512, - Transports: conf.RTSPTransports{gortsplib.ProtocolTCP: {}}, - PathManager: pathManager, - Parent: test.NilLogger, - } - err := s.Initialize() - require.NoError(t, err) - defer s.Close() - - source := gortsplib.Client{} - - media0 := test.UniqueMediaH264() - - err = source.StartRecording( - "rtsp://myuser:mypass@127.0.0.1:8557/teststream?param=value", - &description.Session{Medias: []*description.Media{media0}}) - require.NoError(t, err) - defer source.Close() - - err = source.WritePacketRTP(media0, &rtp.Packet{ - Header: rtp.Header{ - Version: 2, - Marker: true, - PayloadType: 96, - SequenceNumber: 123, - Timestamp: 45343, - SSRC: 563423, - }, - Payload: []byte{5, 2, 3, 4}, - }) - require.NoError(t, err) - - <-dataReceived - - list, err := s.APISessionsList() - require.NoError(t, err) - require.Equal(t, &defs.APIRTSPSessionList{ - Items: []defs.APIRTSPSession{ - { - ID: list.Items[0].ID, - Created: list.Items[0].Created, - RemoteAddr: list.Items[0].RemoteAddr, - State: "publish", - Path: "teststream", - Query: "param=value", - User: "myuser", - UserAgent: list.Items[0].UserAgent, - InboundBytes: list.Items[0].InboundBytes, - InboundRTPPackets: list.Items[0].InboundRTPPackets, - OutboundBytes: list.Items[0].OutboundBytes, - BytesReceived: list.Items[0].BytesReceived, - BytesSent: list.Items[0].BytesSent, - Conns: list.Items[0].Conns, - RTPPacketsReceived: list.Items[0].RTPPacketsReceived, - Transport: new("TCP"), - Profile: new("AVP"), - }, - }, - }, list) - }) + } } } diff --git a/internal/staticsources/mpegts/source.go b/internal/staticsources/mpegts/source.go index a0fef765..18fbb494 100644 --- a/internal/staticsources/mpegts/source.go +++ b/internal/staticsources/mpegts/source.go @@ -58,9 +58,17 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { } if s.DumpPackets { - l.Listen = (&packetdumper.Listen{ - Prefix: "mpegts_source_unix_conn", - }).Do + l.Listen = func(network, address string) (net.Listener, error) { + ln, err2 := net.Listen(network, address) + if err2 != nil { + return nil, err2 + } + + return &packetdumper.Listener{ + Wrapped: ln, + Prefix: "mpegts_source_unix_conn", + }, nil + } } err = l.Initialize() @@ -76,6 +84,7 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { } params := udp.URLToParams(u) + l := &udp.Listener{ Address: params.Address, Source: params.Source, @@ -83,10 +92,27 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { UDPReadBufferSize: int(udpReadBufferSize), } - if s.DumpPackets { - l.ListenPacket = (&packetdumper.ListenPacket{ - Prefix: "mpegts_source_packet_conn", - }).Do + l.ListenPacket = func(network, address string) (net.PacketConn, error) { + pc, err2 := net.ListenPacket(network, address) + if err2 != nil { + return nil, err2 + } + + if s.DumpPackets { + pc2 := &packetdumper.PacketConn{ + Wrapped: pc, + Prefix: "mpegts_source_packet_conn", + } + err2 = pc2.Initialize() + if err2 != nil { + pc.Close() //nolint:errcheck + return nil, err2 + } + + pc = pc2 + } + + return pc, nil } err = l.Initialize() diff --git a/internal/staticsources/rtp/source.go b/internal/staticsources/rtp/source.go index 95bf9bca..c8ba400f 100644 --- a/internal/staticsources/rtp/source.go +++ b/internal/staticsources/rtp/source.go @@ -92,10 +92,27 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { UDPReadBufferSize: int(udpReadBufferSize), } - if s.DumpPackets { - l.ListenPacket = (&packetdumper.ListenPacket{ - Prefix: "rtp_source_packet_conn", - }).Do + l.ListenPacket = func(network, address string) (net.PacketConn, error) { + pc, err2 := net.ListenPacket(network, address) + if err2 != nil { + return nil, err2 + } + + if s.DumpPackets { + pc2 := &packetdumper.PacketConn{ + Wrapped: pc, + Prefix: "rtp_source_packet_conn", + } + err2 = pc2.Initialize() + if err2 != nil { + pc.Close() //nolint:errcheck + return nil, err2 + } + + pc = pc2 + } + + return pc, nil } err = l.Initialize() diff --git a/internal/staticsources/rtsp/source.go b/internal/staticsources/rtsp/source.go index 27f72add..e12dd976 100644 --- a/internal/staticsources/rtsp/source.go +++ b/internal/staticsources/rtsp/source.go @@ -3,6 +3,7 @@ package rtsp import ( "fmt" + "net" "net/url" "time" @@ -187,10 +188,6 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { Prefix: u.Scheme + "_source_conn", }).Do - c.ListenPacket = (&packetdumper.ListenPacket{ - Prefix: u.Scheme + "_source_packet_conn", - }).Do - c.DialTLSContext = (&packetdumper.DialTLSContext{ DialContext: c.DialContext, TLSConfig: tlsConfig, @@ -199,6 +196,29 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error { c.TLSConfig = tlsConfig } + c.ListenPacket = func(network, address string) (net.PacketConn, error) { + pc, err2 := net.ListenPacket(network, address) + if err2 != nil { + return nil, err2 + } + + if s.DumpPackets { + pc2 := &packetdumper.PacketConn{ + Wrapped: pc, + Prefix: u.Scheme + "_source_packet_conn", + } + err2 = pc2.Initialize() + if err2 != nil { + pc.Close() //nolint:errcheck + return nil, err2 + } + + pc = pc2 + } + + return pc, nil + } + err = c.Start() if err != nil { return err diff --git a/mediamtx.yml b/mediamtx.yml index 79a7bb37..25109028 100644 --- a/mediamtx.yml +++ b/mediamtx.yml @@ -281,6 +281,10 @@ rtspServerCert: server.crt # Authentication methods. Available are "basic" and "digest". # "digest" doesn't provide any additional security and is available for compatibility only. rtspAuthMethods: [basic] +# IPs or CIDRs of proxies placed before the RTSP server. +# If the server receives a request from one of these entries, IP in logs +# and authentication will be taken from the PROXY protocol header. +rtspTrustedProxies: [] ############################################### # Global settings -> RTMP server @@ -301,6 +305,10 @@ rtmpsAddress: :1936 rtmpServerKey: server.key # Path to the server certificate. This is needed only when encryption is "strict" or "optional". rtmpServerCert: server.crt +# IPs or CIDRs of proxies placed before the RTMP server. +# If the server receives a request from one of these entries, IP in logs +# and authentication will be taken from the PROXY protocol header. +rtmpTrustedProxies: [] ############################################### # Global settings -> HLS server