Explorar o código

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 hai 13 horas
pai
achega
805f94a00c

+ 1 - 2
internal/amneziawgnet/bench_test.go

@@ -8,7 +8,6 @@ import (
 	"testing"
 	"testing"
 	"time"
 	"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/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/buffer"
 	"gvisor.dev/gvisor/pkg/buffer"
@@ -159,7 +158,7 @@ func newBenchTunnel(b *testing.B, listenPort int, serverAddr, clientAddr string)
 	if err != nil {
 	if err != nil {
 		b.Fatalf("client CreateNetTUN: %v", err)
 		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)
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
 	if err != nil {
 	if err != nil {

+ 3 - 4
internal/amneziawgnet/device_test.go

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

+ 1 - 2
internal/amneziawgnet/diagnostics_test.go

@@ -8,7 +8,6 @@ import (
 	"testing"
 	"testing"
 	"time"
 	"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/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
 	"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
@@ -111,7 +110,7 @@ func TestDiagnoseDeviceReportsListenPortAndPeerState(t *testing.T) {
 	if err != nil {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 		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()
 	defer clientDev.Close()
 
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
 	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) {
 		if raw != "" && !isWildcardListen(raw) {
 			logger.Warningf("amneziawgnet: listen %q is not a bindable IP; using dual-stack wildcard", raw)
 			logger.Warningf("amneziawgnet: listen %q is not a bindable IP; using dual-stack wildcard", raw)
 		}
 		}
-		return awgconn.NewDefaultBind()
+		return wildcardBind()
 	}
 	}
 	if !listenBindable(addr) {
 	if !listenBindable(addr) {
 		logger.Warningf("amneziawgnet: listen %q is not usable on this host; using dual-stack wildcard", raw)
 		logger.Warningf("amneziawgnet: listen %q is not usable on this host; using dual-stack wildcard", raw)
-		return awgconn.NewDefaultBind()
+		return wildcardBind()
 	}
 	}
 	return newPinnedBind(addr)
 	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
 // normalizedListenFP collapses wildcard spellings so fingerprint rebuilds
 // only when the effective Bind actually changes.
 // only when the effective Bind actually changes.
 func normalizedListenFP(listen string) string {
 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) {
 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"} {
 	for _, listen := range []string{"", "0.0.0.0", "::", "::0", "[::]", "hostname.example", "203.0.113.10", "not-an-ip"} {
 		bind := newListenBind(listen)
 		bind := newListenBind(listen)
 		if _, ok := bind.(*pinnedBind); ok {
 		if _, ok := bind.(*pinnedBind); ok {

+ 2 - 2
internal/amneziawgnet/portfwd.go

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

+ 2 - 3
internal/amneziawgnet/portfwd_test.go

@@ -10,7 +10,6 @@ import (
 	"testing"
 	"testing"
 	"time"
 	"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/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip/stack"
 	"gvisor.dev/gvisor/pkg/tcpip/stack"
@@ -190,7 +189,7 @@ func TestPortForwardSetReconcileSurvivesPreBoundPort(t *testing.T) {
 
 
 	const collidingPort = 58911
 	const collidingPort = 58911
 	const okPort = 58912
 	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 {
 	if err != nil {
 		t.Fatalf("pre-bind test port: %v", err)
 		t.Fatalf("pre-bind test port: %v", err)
 	}
 	}
@@ -284,7 +283,7 @@ func TestPortForwardRoundTripTCPAndUDP(t *testing.T) {
 	if err != nil {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 		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()
 	defer clientDev.Close()
 	// clientDev.Close() closes the tun's packet channel without waiting for
 	// clientDev.Close() closes the tun's packet channel without waiting for
 	// writers, so every goroutine writing into clientNet must be gone first.
 	// 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 (
 import (
 	"context"
 	"context"
-	"fmt"
 	"net"
 	"net"
 	"net/netip"
 	"net/netip"
+	"strconv"
 	"sync"
 	"sync"
 	"time"
 	"time"
 
 
@@ -46,7 +46,7 @@ type udpForwardListener struct {
 // toward target(key.email). Bind-failure contract matches
 // toward target(key.email). Bind-failure contract matches
 // listenPortForwardTCP exactly: log, return nil, Reconcile retries later.
 // listenPortForwardTCP exactly: log, return nil, Reconcile retries later.
 func listenPortForwardUDP(gstack *stack.Stack, inboundID int, key portForwardKey, target portForwardTargetFunc) *udpForwardListener {
 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 {
 	if err != nil {
 		logger.Warningf("amneziawgnet: port-forward: inbound %d peer %q: listen udp :%d: %v", inboundID, key.email, key.port, err)
 		logger.Warningf("amneziawgnet: port-forward: inbound %d peer %q: listen udp :%d: %v", inboundID, key.email, key.port, err)
 		return nil
 		return nil

+ 2 - 3
internal/amneziawgnet/relay_e2e_test.go

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

+ 1 - 2
internal/amneziawgnet/udp_test.go

@@ -6,7 +6,6 @@ import (
 	"testing"
 	"testing"
 	"time"
 	"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/device"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
 	"gvisor.dev/gvisor/pkg/tcpip"
 	"gvisor.dev/gvisor/pkg/tcpip"
@@ -98,7 +97,7 @@ func TestNewDeviceUDPHandlerAndReply(t *testing.T) {
 	if err != nil {
 	if err != nil {
 		t.Fatalf("client CreateNetTUN: %v", err)
 		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()
 	defer clientDev.Close()
 
 
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)
 	clientPrivHex, err := wireguard.KeyToHex(clientPriv)