Просмотр исходного кода

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 13 часов назад
Родитель
Сommit
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)