mirror of
https://github.com/netbirdio/ice.git
synced 2026-05-22 17:10:58 -07:00
Add nonblock write option to TCPMux
add nonblock write option for sfu broadcast don't want slow connetion affect other clients.
This commit is contained in:
+9
-3
@@ -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 {
|
||||
|
||||
+57
-39
@@ -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) {
|
||||
|
||||
+67
-3
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user