From 90dc7280e2d22848288a4ce9480b0d2b66d76fb0 Mon Sep 17 00:00:00 2001 From: cnderrauber Date: Wed, 20 Jul 2022 14:06:32 +0800 Subject: [PATCH] Add nonblock write option to TCPMux add nonblock write option for sfu broadcast don't want slow connetion affect other clients. --- tcp_mux.go | 12 ++++-- tcp_mux_test.go | 96 +++++++++++++++++++++++++++------------------- tcp_packet_conn.go | 70 +++++++++++++++++++++++++++++++-- 3 files changed, 133 insertions(+), 45 deletions(-) diff --git a/tcp_mux.go b/tcp_mux.go index ce8f3e1..836c942 100644 --- a/tcp_mux.go +++ b/tcp_mux.go @@ -61,6 +61,11 @@ type TCPMuxParams struct { Listener net.Listener Logger logging.LeveledLogger ReadBufferSize int + + // max buffer size for write op. 0 means no write buffer, the write op will block until the whole packet is written + // if the write buffer is full, the subsequent write packet will be dropped until it has enough space. + // a default 4MB is recommended. + WriteBufferSize int } // NewTCPMuxDefault creates a new instance of TCPMuxDefault. @@ -127,9 +132,10 @@ func (m *TCPMuxDefault) GetConnByUfrag(ufrag string, isIPv6 bool) (net.PacketCon func (m *TCPMuxDefault) createConn(ufrag string, localAddr net.Addr, isIPv6 bool) *tcpPacketConn { conn := newTCPPacketConn(tcpPacketParams{ - ReadBuffer: m.params.ReadBufferSize, - LocalAddr: localAddr, - Logger: m.params.Logger, + ReadBuffer: m.params.ReadBufferSize, + WriteBuffer: m.params.WriteBufferSize, + LocalAddr: localAddr, + Logger: m.params.Logger, }) if isIPv6 { diff --git a/tcp_mux_test.go b/tcp_mux_test.go index 56fee6e..71c6881 100644 --- a/tcp_mux_test.go +++ b/tcp_mux_test.go @@ -18,55 +18,73 @@ var ( ) func TestTCPMux_Recv(t *testing.T) { - report := test.CheckRoutines(t) - defer report() + for name, buffersize := range map[string]int{ + "no buffer": 0, + "buffered 4MB": 4 * 1024 * 1024, + } { + bufSize := buffersize + t.Run(name, func(t *testing.T) { + report := test.CheckRoutines(t) + defer report() - loggerFactory := logging.NewDefaultLoggerFactory() + loggerFactory := logging.NewDefaultLoggerFactory() - listener, err := net.ListenTCP("tcp", &net.TCPAddr{ - IP: net.IP{127, 0, 0, 1}, - Port: 0, - }) - require.NoError(t, err, "error starting listener") - defer func() { - _ = listener.Close() - }() + listener, err := net.ListenTCP("tcp", &net.TCPAddr{ + IP: net.IP{127, 0, 0, 1}, + Port: 0, + }) + require.NoError(t, err, "error starting listener") + defer func() { + _ = listener.Close() + }() - tcpMux := NewTCPMuxDefault(TCPMuxParams{ - Listener: listener, - Logger: loggerFactory.NewLogger("ice"), - ReadBufferSize: 20, - }) + tcpMux := NewTCPMuxDefault(TCPMuxParams{ + Listener: listener, + Logger: loggerFactory.NewLogger("ice"), + ReadBufferSize: 20, + WriteBufferSize: bufSize, + }) - defer func() { - _ = tcpMux.Close() - }() + defer func() { + _ = tcpMux.Close() + }() - require.NotNil(t, tcpMux.LocalAddr(), "tcpMux.LocalAddr() is nil") + require.NotNil(t, tcpMux.LocalAddr(), "tcpMux.LocalAddr() is nil") - conn, err := net.DialTCP("tcp", nil, tcpMux.LocalAddr().(*net.TCPAddr)) - require.NoError(t, err, "error dialing test tcp connection") + conn, err := net.DialTCP("tcp", nil, tcpMux.LocalAddr().(*net.TCPAddr)) + require.NoError(t, err, "error dialing test tcp connection") - msg := stun.New() - msg.Type = stun.MessageType{Method: stun.MethodBinding, Class: stun.ClassRequest} - msg.Add(stun.AttrUsername, []byte("myufrag:otherufrag")) - msg.Encode() + msg := stun.New() + msg.Type = stun.MessageType{Method: stun.MethodBinding, Class: stun.ClassRequest} + msg.Add(stun.AttrUsername, []byte("myufrag:otherufrag")) + msg.Encode() - n, err := writeStreamingPacket(conn, msg.Raw) - require.NoError(t, err, "error writing tcp stun packet") + n, err := writeStreamingPacket(conn, msg.Raw) + require.NoError(t, err, "error writing tcp stun packet") - pktConn, err := tcpMux.GetConnByUfrag("myufrag", false) - require.NoError(t, err, "error retrieving muxed connection for ufrag") - defer func() { - _ = pktConn.Close() - }() + pktConn, err := tcpMux.GetConnByUfrag("myufrag", false) + require.NoError(t, err, "error retrieving muxed connection for ufrag") + defer func() { + _ = pktConn.Close() + }() - recv := make([]byte, n) - n2, raddr, err := pktConn.ReadFrom(recv) - require.NoError(t, err, "error receiving data") - assert.Equal(t, conn.LocalAddr(), raddr, "remote tcp address mismatch") - assert.Equal(t, n, n2, "received byte size mismatch") - assert.Equal(t, msg.Raw, recv, "received bytes mismatch") + recv := make([]byte, n) + n2, raddr, err := pktConn.ReadFrom(recv) + require.NoError(t, err, "error receiving data") + assert.Equal(t, conn.LocalAddr(), raddr, "remote tcp address mismatch") + assert.Equal(t, n, n2, "received byte size mismatch") + assert.Equal(t, msg.Raw, recv, "received bytes mismatch") + + // check echo response + n, err = pktConn.WriteTo(recv, conn.LocalAddr()) + require.NoError(t, err, "error writing echo stun packet") + recvEcho := make([]byte, n) + n3, err := readStreamingPacket(conn, recvEcho) + require.NoError(t, err, "error receiving echo data") + assert.Equal(t, n2, n3, "received byte size mismatch") + assert.Equal(t, msg.Raw, recvEcho, "received bytes mismatch") + }) + } } func TestTCPMux_NoDeadlockWhenClosingUnusedPacketConn(t *testing.T) { diff --git a/tcp_packet_conn.go b/tcp_packet_conn.go index ce3b977..6599759 100644 --- a/tcp_packet_conn.go +++ b/tcp_packet_conn.go @@ -1,15 +1,75 @@ package ice import ( + "errors" "fmt" "io" "net" "sync" + "sync/atomic" "time" "github.com/pion/logging" + "github.com/pion/transport/packetio" ) +type bufferedConn struct { + net.Conn + buffer *packetio.Buffer + logger logging.LeveledLogger + closed int32 +} + +func newBufferedConn(conn net.Conn, bufferSize int, logger logging.LeveledLogger) net.Conn { + buffer := packetio.NewBuffer() + if bufferSize > 0 { + buffer.SetLimitSize(bufferSize) + } + + bc := &bufferedConn{ + Conn: conn, + buffer: buffer, + logger: logger, + } + + go bc.writeProcess() + return bc +} + +func (bc *bufferedConn) Write(b []byte) (int, error) { + n, err := bc.buffer.Write(b) + if err != nil { + return n, err + } + return n, nil +} + +func (bc *bufferedConn) writeProcess() { + pktBuf := make([]byte, receiveMTU) + for atomic.LoadInt32(&bc.closed) == 0 { + n, err := bc.buffer.Read(pktBuf) + if errors.Is(err, io.EOF) { + return + } + + if err != nil { + bc.logger.Warnf("read buffer error: %s", err) + continue + } + + if _, err := bc.Conn.Write(pktBuf[:n]); err != nil { + bc.logger.Warnf("write error: %s", err) + continue + } + } +} + +func (bc *bufferedConn) Close() error { + atomic.StoreInt32(&bc.closed, 1) + _ = bc.buffer.Close() + return bc.Conn.Close() +} + type tcpPacketConn struct { params *tcpPacketParams @@ -31,9 +91,10 @@ type streamingPacket struct { } type tcpPacketParams struct { - ReadBuffer int - LocalAddr net.Addr - Logger logging.LeveledLogger + ReadBuffer int + LocalAddr net.Addr + Logger logging.LeveledLogger + WriteBuffer int } func newTCPPacketConn(params tcpPacketParams) *tcpPacketConn { @@ -65,6 +126,9 @@ func (t *tcpPacketConn) AddConn(conn net.Conn, firstPacketData []byte) error { return fmt.Errorf("%w: %s", errConnectionAddrAlreadyExist, conn.RemoteAddr().String()) } + if t.params.WriteBuffer > 0 { + conn = newBufferedConn(conn, t.params.WriteBuffer, t.params.Logger) + } t.conns[conn.RemoteAddr().String()] = conn t.wg.Add(1)