diff --git a/pkg/proto/udp/udp.go b/pkg/proto/udp/udp.go index 70b51900..ec3bd94f 100644 --- a/pkg/proto/udp/udp.go +++ b/pkg/proto/udp/udp.go @@ -63,9 +63,13 @@ func ForwardUserConn(udpConn *net.UDPConn, readCh <-chan *msg.UDPPacket, sendCh // NewUDPPacket copies buf[:n], so the read buffer can be reused udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr) - select { - case sendCh <- udpMsg: - default: + if err = errors.PanicToError(func() { + select { + case sendCh <- udpMsg: + default: + } + }); err != nil { + return } } } diff --git a/pkg/proto/udp/udp_test.go b/pkg/proto/udp/udp_test.go index 1a7f0091..b959a26a 100644 --- a/pkg/proto/udp/udp_test.go +++ b/pkg/proto/udp/udp_test.go @@ -1,9 +1,13 @@ package udp import ( + "net" "testing" + "time" "github.com/stretchr/testify/require" + + "github.com/fatedier/frp/pkg/msg" ) func TestUdpPacket(t *testing.T) { @@ -16,3 +20,33 @@ func TestUdpPacket(t *testing.T) { require.NoError(err) require.EqualValues(buf, newBuf) } + +func TestForwardUserConnReturnsWhenSendChannelIsClosed(t *testing.T) { + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + + readCh := make(chan *msg.UDPPacket) + sendCh := make(chan *msg.UDPPacket) + close(sendCh) + t.Cleanup(func() { close(readCh) }) + + done := make(chan struct{}) + go func() { + ForwardUserConn(listener, readCh, sendCh, 1500) + close(done) + }() + + sender, err := net.DialUDP("udp4", nil, listener.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + t.Cleanup(func() { _ = sender.Close() }) + + _, err = sender.Write([]byte("trigger")) + require.NoError(t, err) + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("ForwardUserConn did not return after sending to a closed channel") + } +}