|
@@ -6,6 +6,7 @@ import (
|
|
|
"io"
|
|
"io"
|
|
|
"net"
|
|
"net"
|
|
|
"net/netip"
|
|
"net/netip"
|
|
|
|
|
+ "sync"
|
|
|
"testing"
|
|
"testing"
|
|
|
"time"
|
|
"time"
|
|
|
|
|
|
|
@@ -285,6 +286,10 @@ func TestPortForwardRoundTripTCPAndUDP(t *testing.T) {
|
|
|
}
|
|
}
|
|
|
clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
|
|
clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
|
|
|
defer clientDev.Close()
|
|
defer clientDev.Close()
|
|
|
|
|
+ // clientDev.Close() closes the tun's packet channel without waiting for
|
|
|
|
|
+ // writers, so every goroutine writing into clientNet must be gone first.
|
|
|
|
|
+ var clientSvc sync.WaitGroup
|
|
|
|
|
+ defer clientSvc.Wait()
|
|
|
|
|
|
|
|
clientPrivHex, err := wireguard.KeyToHex(clientPriv)
|
|
clientPrivHex, err := wireguard.KeyToHex(clientPriv)
|
|
|
if err != nil {
|
|
if err != nil {
|
|
@@ -338,13 +343,16 @@ primed:
|
|
|
t.Fatalf("client ListenTCP: %v", err)
|
|
t.Fatalf("client ListenTCP: %v", err)
|
|
|
}
|
|
}
|
|
|
defer tcpSvc.Close()
|
|
defer tcpSvc.Close()
|
|
|
|
|
+ clientSvc.Add(1)
|
|
|
go func() {
|
|
go func() {
|
|
|
|
|
+ defer clientSvc.Done()
|
|
|
for {
|
|
for {
|
|
|
c, err := tcpSvc.Accept()
|
|
c, err := tcpSvc.Accept()
|
|
|
if err != nil {
|
|
if err != nil {
|
|
|
return
|
|
return
|
|
|
}
|
|
}
|
|
|
- go func() { io.Copy(c, c); c.Close() }()
|
|
|
|
|
|
|
+ clientSvc.Add(1)
|
|
|
|
|
+ go func() { defer clientSvc.Done(); io.Copy(c, c); c.Close() }()
|
|
|
}
|
|
}
|
|
|
}()
|
|
}()
|
|
|
|
|
|
|
@@ -353,7 +361,9 @@ primed:
|
|
|
t.Fatalf("client ListenUDP: %v", err)
|
|
t.Fatalf("client ListenUDP: %v", err)
|
|
|
}
|
|
}
|
|
|
defer udpSvc.Close()
|
|
defer udpSvc.Close()
|
|
|
|
|
+ clientSvc.Add(1)
|
|
|
go func() {
|
|
go func() {
|
|
|
|
|
+ defer clientSvc.Done()
|
|
|
buf := make([]byte, 1500)
|
|
buf := make([]byte, 1500)
|
|
|
for {
|
|
for {
|
|
|
n, addr, err := udpSvc.ReadFrom(buf)
|
|
n, addr, err := udpSvc.ReadFrom(buf)
|