Parcourir la source

chore(amneziawgnet): bind test sockets to loopback

Windows Firewall prompted on every run of amneziawgnet.test.exe, since
the tests opened AmneziaWG, outbound-client and port-forward sockets on
all interfaces and go test rebuilds the binary under a fresh temp path,
so a granted exception never sticks.

Route every "all interfaces" bind through wildcardBindHost, which is
empty in production (behaviour unchanged) and pinned to 127.0.0.1 by
the package TestMain.
MHSanaei il y a 14 heures
Parent
commit
805f94a00c

+ 1 - 2
internal/amneziawgnet/bench_test.go

@@ -8,7 +8,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/buffer"
@@ -159,7 +158,7 @@ func newBenchTunnel(b *testing.B, listenPort int, serverAddr, clientAddr string)
 	if err != nil {
 		b.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
 	if err != nil {

+ 3 - 4
internal/amneziawgnet/device_test.go

@@ -10,7 +10,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
@@ -107,7 +106,7 @@ func TestNewDeviceHandshakeForwarderAndIdentity(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
@@ -314,7 +313,7 @@ func TestNewDeviceHeaderProtectionAndContentPaddingRoundTrip(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
@@ -492,7 +491,7 @@ func TestNewDeviceRandomTrailersAndDisableCookiesRoundTrip(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)

+ 1 - 2
internal/amneziawgnet/diagnostics_test.go

@@ -8,7 +8,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
@@ -111,7 +110,7 @@ func TestDiagnoseDeviceReportsListenPortAndPeerState(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)

+ 11 - 0
internal/amneziawgnet/main_test.go

@@ -0,0 +1,11 @@
+package amneziawgnet
+
+import (
+	"os"
+	"testing"
+)
+
+func TestMain(m *testing.M) {
+	wildcardBindHost = "127.0.0.1"
+	os.Exit(m.Run())
+}

+ 13 - 2
internal/amneziawgnet/pinned_bind.go

@@ -162,15 +162,26 @@ func newListenBind(listen string) awgconn.Bind {
 		if raw != "" && !isWildcardListen(raw) {
 			logger.Warningf("amneziawgnet: listen %q is not a bindable IP; using dual-stack wildcard", raw)
 		}
-		return awgconn.NewDefaultBind()
+		return wildcardBind()
 	}
 	if !listenBindable(addr) {
 		logger.Warningf("amneziawgnet: listen %q is not usable on this host; using dual-stack wildcard", raw)
-		return awgconn.NewDefaultBind()
+		return wildcardBind()
 	}
 	return newPinnedBind(addr)
 }
 
+// wildcardBindHost replaces "all interfaces" for every host socket this package
+// opens; TestMain pins it to loopback so Windows Firewall never prompts.
+var wildcardBindHost = ""
+
+func wildcardBind() awgconn.Bind {
+	if wildcardBindHost != "" {
+		return newPinnedBind(netip.MustParseAddr(wildcardBindHost))
+	}
+	return awgconn.NewDefaultBind()
+}
+
 // normalizedListenFP collapses wildcard spellings so fingerprint rebuilds
 // only when the effective Bind actually changes.
 func normalizedListenFP(listen string) string {

+ 3 - 0
internal/amneziawgnet/pinned_bind_test.go

@@ -74,6 +74,9 @@ func TestNewListenBindPinsSpecificAddress(t *testing.T) {
 }
 
 func TestNewListenBindWildcardUsesDefault(t *testing.T) {
+	prev := wildcardBindHost
+	wildcardBindHost = ""
+	t.Cleanup(func() { wildcardBindHost = prev })
 	for _, listen := range []string{"", "0.0.0.0", "::", "::0", "[::]", "hostname.example", "203.0.113.10", "not-an-ip"} {
 		bind := newListenBind(listen)
 		if _, ok := bind.(*pinnedBind); ok {

+ 2 - 2
internal/amneziawgnet/portfwd.go

@@ -23,9 +23,9 @@ package amneziawgnet
 
 import (
 	"context"
-	"fmt"
 	"net"
 	"net/netip"
+	"strconv"
 	"sync"
 	"time"
 
@@ -304,7 +304,7 @@ type tcpForwardListener struct {
 // result as "not open this round" and retries on every future Reconcile
 // call for as long as the key stays desired.
 func listenPortForwardTCP(gstack *stack.Stack, inboundID int, key portForwardKey, target portForwardTargetFunc) *tcpForwardListener {
-	ln, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", fmt.Sprintf(":%d", key.port))
+	ln, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", net.JoinHostPort(wildcardBindHost, strconv.Itoa(key.port)))
 	if err != nil {
 		logger.Warningf("amneziawgnet: port-forward: inbound %d peer %q: listen tcp :%d: %v", inboundID, key.email, key.port, err)
 		return nil

+ 2 - 3
internal/amneziawgnet/portfwd_test.go

@@ -10,7 +10,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/stack"
@@ -190,7 +189,7 @@ func TestPortForwardSetReconcileSurvivesPreBoundPort(t *testing.T) {
 
 	const collidingPort = 58911
 	const okPort = 58912
-	blocker, err := net.Listen("tcp", fmt.Sprintf(":%d", collidingPort))
+	blocker, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", collidingPort))
 	if err != nil {
 		t.Fatalf("pre-bind test port: %v", err)
 	}
@@ -284,7 +283,7 @@ func TestPortForwardRoundTripTCPAndUDP(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	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.

+ 2 - 2
internal/amneziawgnet/portfwd_udp.go

@@ -2,9 +2,9 @@ package amneziawgnet
 
 import (
 	"context"
-	"fmt"
 	"net"
 	"net/netip"
+	"strconv"
 	"sync"
 	"time"
 
@@ -46,7 +46,7 @@ type udpForwardListener struct {
 // toward target(key.email). Bind-failure contract matches
 // listenPortForwardTCP exactly: log, return nil, Reconcile retries later.
 func listenPortForwardUDP(gstack *stack.Stack, inboundID int, key portForwardKey, target portForwardTargetFunc) *udpForwardListener {
-	pc, err := (&net.ListenConfig{}).ListenPacket(context.Background(), "udp", fmt.Sprintf(":%d", key.port))
+	pc, err := (&net.ListenConfig{}).ListenPacket(context.Background(), "udp", net.JoinHostPort(wildcardBindHost, strconv.Itoa(key.port)))
 	if err != nil {
 		logger.Warningf("amneziawgnet: port-forward: inbound %d peer %q: listen udp :%d: %v", inboundID, key.email, key.port, err)
 		return nil

+ 2 - 3
internal/amneziawgnet/relay_e2e_test.go

@@ -13,7 +13,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
@@ -196,7 +195,7 @@ func TestSocksRelayAgainstRealXray(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
@@ -421,7 +420,7 @@ func TestManagerEnsureAutomaticallyWiresRelay(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)

+ 1 - 2
internal/amneziawgnet/udp_test.go

@@ -6,7 +6,6 @@ import (
 	"testing"
 	"time"
 
-	awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
 	"github.com/amnezia-vpn/amneziawg-go/v3/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip"
@@ -98,7 +97,7 @@ func TestNewDeviceUDPHandlerAndReply(t *testing.T) {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 	}
-	clientDev := device.NewDevice(clientTun, awgconn.NewDefaultBind(), device.NewLogger(device.LogLevelSilent, ""))
+	clientDev := device.NewDevice(clientTun, newListenBind(""), device.NewLogger(device.LogLevelSilent, ""))
 	defer clientDev.Close()
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)