From cefb61e85f31f1eaee65eead5bb608424ca94803 Mon Sep 17 00:00:00 2001 From: Neil Date: Mon, 13 Jul 2026 09:55:01 -0600 Subject: [PATCH 1/3] feat(firewall): scoped temporary bootstrap connection allowances Add a mechanism to open narrowly-scoped, temporary outbound firewall allowances (destination IP + protocol + port + firewall mark, on the physical interface) that are needed before the VPN tunnel is up, then reliably torn down afterwards. - temporary.go: central registry of temporary allowances with retain-on-failure cleanup, idempotent deletion under a bounded non-cancelable context, and a sweep on firewall shutdown and before re-adding on reconnect. Attempted interfaces are tracked before the iptables append so a rule applied during a cancellation race is never left untracked. - iptables: emit and parse the '-m mark --mark ' match so the temporary ACCEPT rules are scoped to the bootstrap socket mark and can be deleted symmetrically. Reject malformed mark match forms instead of silently accepting a zero mark or panicking on a bad value. Covered by unit tests for rollback, deletion-failure retry, idempotence, mark scoping, reconnect de-duplication and mark parsing edge cases. --- internal/firewall/enable.go | 10 +- internal/firewall/firewall.go | 1 + internal/firewall/interfaces.go | 2 + internal/firewall/iptables/iptables.go | 24 ++ internal/firewall/iptables/parse.go | 18 +- internal/firewall/iptables/parse_test.go | 29 ++ internal/firewall/iptables/temporary_test.go | 24 ++ internal/firewall/temporary.go | 158 ++++++++++ internal/firewall/temporary_test.go | 298 +++++++++++++++++++ 9 files changed, 555 insertions(+), 9 deletions(-) create mode 100644 internal/firewall/iptables/temporary_test.go create mode 100644 internal/firewall/temporary.go create mode 100644 internal/firewall/temporary_test.go diff --git a/internal/firewall/enable.go b/internal/firewall/enable.go index 6f5cb18bf..8af8561cf 100644 --- a/internal/firewall/enable.go +++ b/internal/firewall/enable.go @@ -22,7 +22,15 @@ func (c *Config) SetEnabled(ctx context.Context, enabled bool) (err error) { if !enabled { c.logger.Info("disabling...") - c.restore(ctx) + cleanupCtx, cancelCleanup := newTemporaryCleanupContext(ctx) + cleanupErr := c.sweepTemporaryConnectionRules(cleanupCtx, true) + cancelCleanup() + if cleanupErr != nil { + c.logger.Warn(cleanupErr.Error()) + } + restoreCtx, cancelRestore := newTemporaryCleanupContext(ctx) + c.restore(restoreCtx) + cancelRestore() c.enabled = false c.logger.Info("disabled successfully") return nil diff --git a/internal/firewall/firewall.go b/internal/firewall/firewall.go index 6b542a9ec..caddae447 100644 --- a/internal/firewall/firewall.go +++ b/internal/firewall/firewall.go @@ -29,6 +29,7 @@ type Config struct { outboundSubnets []netip.Prefix allowedInputPorts map[uint16]map[string]struct{} // port to interfaces set mapping portRedirections portRedirections + temporaryRules map[*temporaryConnectionRules]struct{} stateMutex sync.Mutex } diff --git a/internal/firewall/interfaces.go b/internal/firewall/interfaces.go index d9c830da2..ec7224eae 100644 --- a/internal/firewall/interfaces.go +++ b/internal/firewall/interfaces.go @@ -28,6 +28,8 @@ type firewallImpl interface { //nolint:interfacebloat AcceptIpv6MulticastOutput(ctx context.Context, intf string) error AcceptOutput(ctx context.Context, protocol, intf string, ip netip.Addr, port uint16, remove bool) error + AcceptOutputMarked(ctx context.Context, protocol, intf string, + ip netip.Addr, port uint16, mark uint32, remove bool) error AcceptOutputFromIPToSubnet(ctx context.Context, intf string, assignedIP netip.Addr, subnet netip.Prefix, remove bool) error AcceptOutputThroughInterface(ctx context.Context, intf string, remove bool) error diff --git a/internal/firewall/iptables/iptables.go b/internal/firewall/iptables/iptables.go index c48879295..d0721b1f2 100644 --- a/internal/firewall/iptables/iptables.go +++ b/internal/firewall/iptables/iptables.go @@ -177,6 +177,30 @@ func (c *Config) AcceptOutput(ctx context.Context, return c.runIP6tablesInstruction(ctx, instruction) } +func (c *Config) AcceptOutputMarked(ctx context.Context, + protocol, intf string, ip netip.Addr, port uint16, mark uint32, remove bool, +) error { + instruction := acceptOutputMarkedInstruction(protocol, intf, ip, port, mark, remove) + if ip.Is4() { + return c.runIptablesInstruction(ctx, instruction) + } else if c.ip6Tables == "" { + return fmt.Errorf("accept marked output to VPN server %s: %s", ip, needIP6Tables) + } + return c.runIP6tablesInstruction(ctx, instruction) +} + +func acceptOutputMarkedInstruction(protocol, intf string, ip netip.Addr, + port uint16, mark uint32, remove bool, +) string { + interfaceFlag := "-o " + intf + if intf == "*" { // all interfaces + interfaceFlag = "" + } + + return fmt.Sprintf("%s OUTPUT -d %s %s -p %s -m %s --dport %d -m mark --mark %d -j ACCEPT", + appendOrDelete(remove), ip, interfaceFlag, protocol, protocol, port, mark) +} + // AcceptOutputFromIPToSubnet accepts outgoing traffic from sourceIP to destinationSubnet // on the interface intf. If intf is empty, it is set to "*" which means all interfaces. // If remove is true, the rule is removed instead of added. diff --git a/internal/firewall/iptables/parse.go b/internal/firewall/iptables/parse.go index 929af05c5..378d1c93d 100644 --- a/internal/firewall/iptables/parse.go +++ b/internal/firewall/iptables/parse.go @@ -128,11 +128,11 @@ func parseInstructionFlag(fields []string, instruction *iptablesInstruction) (co case "--mark": const base = 0 // auto-detect const bits = 32 - value, err := strconv.ParseUint(value, base, bits) + markValue, err := strconv.ParseUint(value, base, bits) if err != nil { - return 0, fmt.Errorf("parsing mark value %q: %w", fields[2], err) + return 0, fmt.Errorf("parsing mark value %q: %w", value, err) } - instruction.mark.value = uint(value) + instruction.mark.value = uint(markValue) case "-i", "--in-interface": instruction.inputInterface = value case "-o", "--out-interface": @@ -232,13 +232,15 @@ func parseMatchModule(fields []string, instruction *iptablesInstruction) ( // parse it twice. case "mark": consumed++ - switch fields[consumed] { - case "!": + // An optional "!" negates the match. The "--mark " flag that + // follows (whether negated or not) is parsed by the main flag loop, + // so we must not consume it here. + if consumed < len(fields) && fields[consumed] == "!" { consumed++ instruction.mark.invert = true - default: - return consumed, fmt.Errorf("iptables command is malformed: unsupported match mark with value: %s", - fields[2]) + } + if consumed+1 >= len(fields) || fields[consumed] != "--mark" { + return 0, errors.New("iptables command is malformed: mark match requires --mark followed by a value") } default: return 0, fmt.Errorf("iptables command is malformed: unknown match value: %s", diff --git a/internal/firewall/iptables/parse_test.go b/internal/firewall/iptables/parse_test.go index 0928c2f4b..b06203547 100644 --- a/internal/firewall/iptables/parse_test.go +++ b/internal/firewall/iptables/parse_test.go @@ -26,6 +26,35 @@ func Test_parseIptablesInstruction(t *testing.T) { s: "-x something", errMessage: "parsing \"-x something\": iptables command is malformed: unknown key \"-x\"", }, + "mark_match_missing_flag_and_value": { + s: "-m mark", + errMessage: "parsing \"-m mark\": parsing match module: " + + "iptables command is malformed: mark match requires --mark followed by a value", + }, + "inverted_mark_match_missing_flag_and_value": { + s: "-m mark !", + errMessage: "parsing \"-m mark !\": parsing match module: " + + "iptables command is malformed: mark match requires --mark followed by a value", + }, + "invalid_mark_value": { + s: "--mark bogus", + errMessage: "parsing \"--mark bogus\": parsing mark value \"bogus\": " + + "strconv.ParseUint: parsing \"bogus\": invalid syntax", + }, + "mark_match": { + s: "-m mark --mark 51820", + instruction: iptablesInstruction{ + table: "filter", + mark: mark{value: 51820}, + }, + }, + "inverted_mark_match": { + s: "-m mark ! --mark 51820", + instruction: iptablesInstruction{ + table: "filter", + mark: mark{value: 51820, invert: true}, + }, + }, "one_pair": { s: "-A INPUT", instruction: iptablesInstruction{ diff --git a/internal/firewall/iptables/temporary_test.go b/internal/firewall/iptables/temporary_test.go new file mode 100644 index 000000000..f2fa7aaef --- /dev/null +++ b/internal/firewall/iptables/temporary_test.go @@ -0,0 +1,24 @@ +package iptables + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_acceptOutputMarkedInstruction_requiresBootstrapSocketMark(t *testing.T) { + t.Parallel() + + const bootstrapFirewallMark uint32 = 51820 + instruction := acceptOutputMarkedInstruction("tcp", "eth0", + netip.MustParseAddr("198.51.100.10"), 443, bootstrapFirewallMark, false) + + assert.Equal(t, "--append OUTPUT -d 198.51.100.10 -o eth0 -p tcp -m tcp --dport 443 "+ + "-m mark --mark 51820 -j ACCEPT", instruction) + parsedInstruction, err := parseIptablesInstruction(instruction) + require.NoError(t, err) + assert.Equal(t, uint(bootstrapFirewallMark), parsedInstruction.mark.value) + assert.False(t, parsedInstruction.mark.invert) +} diff --git a/internal/firewall/temporary.go b/internal/firewall/temporary.go new file mode 100644 index 000000000..67b64bcc6 --- /dev/null +++ b/internal/firewall/temporary.go @@ -0,0 +1,158 @@ +package firewall + +import ( + "context" + "errors" + "fmt" + "net/netip" + "time" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/models" + "github.com/qdm12/gluetun/internal/netlink" + "github.com/qdm12/gluetun/internal/routing" +) + +type temporaryConnectionRules struct { + connection models.Connection + interfaces []string + cleanupRequested bool +} + +// TempAllowConnection temporarily allows one exact destination IP, protocol +// and port through the default-route interfaces matching its address family. +// The returned function removes every rule added by this call. +func (c *Config) TempAllowConnection(ctx context.Context, connection models.Connection) ( + remove func(context.Context) error, err error, +) { + switch { + case !connection.IP.IsValid(): + return nil, errors.New("connection IP is not set") + case connection.Port == 0: + return nil, errors.New("connection port is not set") + case connection.Protocol != constants.TCP && connection.Protocol != constants.UDP: + return nil, fmt.Errorf("connection protocol is not supported: %s", connection.Protocol) + } + + c.stateMutex.Lock() + defer c.stateMutex.Unlock() + + if !c.enabled { + return func(context.Context) error { return nil }, nil + } + + cleanupCtx, cancelCleanup := newTemporaryCleanupContext(ctx) + err = c.sweepTemporaryConnectionRules(cleanupCtx, false) + cancelCleanup() + if err != nil { + return nil, fmt.Errorf("sweeping previous temporary output connections: %w", err) + } + + interfaces := defaultRouteInterfacesForIP(c.defaultRoutes, connection.IP) + if len(interfaces) == 0 { + return nil, errors.New("default route for connection IP address family is not found") + } + + rules := &temporaryConnectionRules{ + connection: connection, + interfaces: make([]string, 0, len(interfaces)), + } + for _, interfaceName := range interfaces { + rules.interfaces = append(rules.interfaces, interfaceName) + if c.temporaryRules == nil { + c.temporaryRules = make(map[*temporaryConnectionRules]struct{}) + } + c.temporaryRules[rules] = struct{}{} + + const bootstrapFirewallMark uint32 = 51820 + const remove = false + err = c.impl.AcceptOutputMarked(ctx, connection.Protocol, interfaceName, + connection.IP, connection.Port, bootstrapFirewallMark, remove) + if err != nil { + rules.cleanupRequested = true + cleanupCtx, cancelCleanup := newTemporaryCleanupContext(ctx) + cleanupErr := c.removeTemporaryConnectionRules(cleanupCtx, rules) + cancelCleanup() + allowErr := fmt.Errorf("allowing temporary output connection: %w", err) + return nil, errors.Join(allowErr, cleanupErr) + } + } + + return func(removeCtx context.Context) error { + c.stateMutex.Lock() + defer c.stateMutex.Unlock() + _, outstanding := c.temporaryRules[rules] + if !outstanding { + return nil + } + + rules.cleanupRequested = true + cleanupCtx, cancelCleanup := newTemporaryCleanupContext(removeCtx) + defer cancelCleanup() + return c.removeTemporaryConnectionRules(cleanupCtx, rules) + }, nil +} + +func (c *Config) removeTemporaryConnectionRules(ctx context.Context, rules *temporaryConnectionRules) error { + remainingInterfaces := rules.interfaces[:0] + errs := make([]error, 0, len(rules.interfaces)) + for _, interfaceName := range rules.interfaces { + const bootstrapFirewallMark uint32 = 51820 + const remove = true + err := c.impl.AcceptOutputMarked(ctx, rules.connection.Protocol, interfaceName, + rules.connection.IP, rules.connection.Port, bootstrapFirewallMark, remove) + if err != nil { + remainingInterfaces = append(remainingInterfaces, interfaceName) + errs = append(errs, err) + } + } + rules.interfaces = remainingInterfaces + if len(rules.interfaces) == 0 { + delete(c.temporaryRules, rules) + } + if len(errs) > 0 { + return fmt.Errorf("removing temporary output connection: %w", errors.Join(errs...)) + } + return nil +} + +func (c *Config) sweepTemporaryConnectionRules(ctx context.Context, includeActive bool) error { + errs := make([]error, 0, len(c.temporaryRules)) + for rules := range c.temporaryRules { + if !includeActive && !rules.cleanupRequested { + continue + } + rules.cleanupRequested = true + err := c.removeTemporaryConnectionRules(ctx, rules) + if err != nil { + errs = append(errs, err) + } + } + if len(errs) > 0 { + return fmt.Errorf("sweeping temporary output connections: %w", errors.Join(errs...)) + } + return nil +} + +func newTemporaryCleanupContext(ctx context.Context) (context.Context, context.CancelFunc) { + const cleanupTimeout = 5 * time.Second + return context.WithTimeout(context.WithoutCancel(ctx), cleanupTimeout) +} + +func defaultRouteInterfacesForIP(defaultRoutes []routing.DefaultRoute, ip netip.Addr) []string { + interfaces := make([]string, 0, len(defaultRoutes)) + seen := make(map[string]struct{}, len(defaultRoutes)) + for _, defaultRoute := range defaultRoutes { + addressFamilyMatches := ip.Is4() == (defaultRoute.Family == netlink.FamilyV4) + if !addressFamilyMatches { + continue + } + _, alreadySeen := seen[defaultRoute.NetInterface] + if alreadySeen { + continue + } + seen[defaultRoute.NetInterface] = struct{}{} + interfaces = append(interfaces, defaultRoute.NetInterface) + } + return interfaces +} diff --git a/internal/firewall/temporary_test.go b/internal/firewall/temporary_test.go new file mode 100644 index 000000000..7ad07c1d8 --- /dev/null +++ b/internal/firewall/temporary_test.go @@ -0,0 +1,298 @@ +package firewall + +import ( + "context" + "errors" + "net/netip" + "testing" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/models" + "github.com/qdm12/gluetun/internal/netlink" + "github.com/qdm12/gluetun/internal/routing" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type temporaryRuleCall struct { + interfaceName string + mark uint32 + remove bool + contextErr error +} + +type temporaryFirewallImpl struct { + firewallImpl + addErrorInterface string + addErrorAfterApplyInterface string + onAddErrorAfterApply func() + removeFailures uint + calls []temporaryRuleCall + activeRules map[string]uint + maxActiveRules uint +} + +type temporaryLogger struct{} + +func (temporaryLogger) Debug(string) {} +func (temporaryLogger) Info(string) {} +func (temporaryLogger) Warn(string) {} +func (temporaryLogger) Error(string) {} + +func (f *temporaryFirewallImpl) AcceptOutputMarked(ctx context.Context, _, interfaceName string, + _ netip.Addr, _ uint16, mark uint32, remove bool, +) error { + f.calls = append(f.calls, temporaryRuleCall{ + interfaceName: interfaceName, + mark: mark, + remove: remove, + contextErr: ctx.Err(), + }) + if !remove && interfaceName == f.addErrorInterface { + return errors.New("add error") + } + if remove && f.removeFailures > 0 { + f.removeFailures-- + return errors.New("remove error") + } + if f.activeRules == nil { + f.activeRules = make(map[string]uint) + } + if remove { + if f.activeRules[interfaceName] <= 1 { + delete(f.activeRules, interfaceName) + } else { + f.activeRules[interfaceName]-- + } + return nil + } + + f.activeRules[interfaceName]++ + var activeRules uint + for _, count := range f.activeRules { + activeRules += count + } + if activeRules > f.maxActiveRules { + f.maxActiveRules = activeRules + } + if interfaceName == f.addErrorAfterApplyInterface { + f.onAddErrorAfterApply() + return errors.New("add error after apply") + } + return nil +} + +func newTemporaryTestConfig(impl firewallImpl, interfaces ...string) *Config { + defaultRoutes := make([]routing.DefaultRoute, len(interfaces)) + for i, interfaceName := range interfaces { + defaultRoutes[i] = routing.DefaultRoute{ + NetInterface: interfaceName, + Family: netlink.FamilyV4, + } + } + return &Config{ + impl: impl, + enabled: true, + defaultRoutes: defaultRoutes, + } +} + +func temporaryTestConnection() models.Connection { + return models.Connection{ + IP: netip.MustParseAddr("198.51.100.10"), + Port: 443, + Protocol: constants.TCP, + } +} + +func Test_Config_TempAllowConnection_partialAdditionRollback(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{addErrorInterface: "eth1"} + config := newTemporaryTestConfig(impl, "eth0", "eth1") + + remove, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + + require.Error(t, err) + assert.Nil(t, remove) + assert.ErrorContains(t, err, "allowing temporary output connection: add error") + assert.Empty(t, impl.activeRules) + assert.Empty(t, config.temporaryRules) + require.Len(t, impl.calls, 4) + assert.Equal(t, []temporaryRuleCall{ + {interfaceName: "eth0", mark: 51820}, + {interfaceName: "eth1", mark: 51820}, + {interfaceName: "eth0", mark: 51820, remove: true}, + {interfaceName: "eth1", mark: 51820, remove: true}, + }, impl.calls) +} + +func Test_Config_TempAllowConnection_appendErrorAfterApplyIsTrackedAndCleaned(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(t.Context()) + var config *Config + trackedBeforeError := false + impl := &temporaryFirewallImpl{ + addErrorAfterApplyInterface: "eth0", + onAddErrorAfterApply: func() { + cancel() + require.Len(t, config.temporaryRules, 1) + for rules := range config.temporaryRules { + assert.Equal(t, []string{"eth0"}, rules.interfaces) + trackedBeforeError = true + } + }, + } + config = newTemporaryTestConfig(impl, "eth0") + + remove, err := config.TempAllowConnection(ctx, temporaryTestConnection()) + + require.Error(t, err) + assert.Nil(t, remove) + assert.ErrorContains(t, err, "allowing temporary output connection: add error after apply") + assert.ErrorIs(t, ctx.Err(), context.Canceled) + assert.True(t, trackedBeforeError) + assert.Empty(t, impl.activeRules) + assert.Empty(t, config.temporaryRules) + require.Len(t, impl.calls, 2) + assert.Equal(t, []temporaryRuleCall{ + {interfaceName: "eth0", mark: 51820}, + {interfaceName: "eth0", mark: 51820, remove: true}, + }, impl.calls) + assert.NoError(t, impl.calls[1].contextErr) +} + +func Test_Config_TempAllowConnection_deletionFailureCanBeRetried(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{removeFailures: 1} + config := newTemporaryTestConfig(impl, "eth0", "eth1") + remove, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + require.NoError(t, err) + + err = remove(t.Context()) + require.Error(t, err) + assert.ErrorContains(t, err, "remove error") + assert.Equal(t, map[string]uint{"eth0": 1}, impl.activeRules) + assert.Len(t, config.temporaryRules, 1) + + require.NoError(t, remove(t.Context())) + assert.Empty(t, impl.activeRules) + assert.Empty(t, config.temporaryRules) + require.Len(t, impl.calls, 5) + assert.Equal(t, temporaryRuleCall{ + interfaceName: "eth0", + mark: 51820, + remove: true, + }, impl.calls[4]) + callCount := len(impl.calls) + require.NoError(t, remove(t.Context())) + assert.Len(t, impl.calls, callCount) +} + +func Test_Config_TempAllowConnection_cancellationStillTearsDown(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{} + config := newTemporaryTestConfig(impl, "eth0") + ctx, cancel := context.WithCancel(t.Context()) + remove, err := config.TempAllowConnection(ctx, temporaryTestConnection()) + require.NoError(t, err) + cancel() + + require.NoError(t, remove(ctx)) + assert.Empty(t, impl.activeRules) + require.Len(t, impl.calls, 2) + assert.NoError(t, impl.calls[1].contextErr) +} + +func Test_Config_TempAllowConnection_reconnectSweepsFailedRulesBeforeAdding(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{removeFailures: 1} + config := newTemporaryTestConfig(impl, "eth0") + firstRemove, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + require.NoError(t, err) + require.Error(t, firstRemove(t.Context())) + + secondRemove, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + require.NoError(t, err) + assert.Equal(t, uint(1), impl.maxActiveRules) + assert.Equal(t, map[string]uint{"eth0": 1}, impl.activeRules) + require.NoError(t, secondRemove(t.Context())) + assert.Empty(t, impl.activeRules) + assert.Empty(t, config.temporaryRules) +} + +func Test_Config_TempAllowConnection_onlyMarkedTrafficIsAllowed(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{} + config := newTemporaryTestConfig(impl, "eth0") + remove, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, remove(t.Context())) }) + + require.Len(t, impl.calls, 1) + assert.Equal(t, uint32(51820), impl.calls[0].mark) + assert.NotEqual(t, uint32(0), impl.calls[0].mark, + "unmarked traffic to the same destination must not match the allowance") +} + +func Test_Config_SetEnabled_shutdownSweepsTemporaryRulesWithCanceledContext(t *testing.T) { + t.Parallel() + + impl := &temporaryFirewallImpl{} + config := newTemporaryTestConfig(impl, "eth0") + config.logger = temporaryLogger{} + restoreContextErr := errors.New("restore was not called") + config.restore = func(ctx context.Context) { + restoreContextErr = ctx.Err() + } + _, err := config.TempAllowConnection(t.Context(), temporaryTestConnection()) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + require.NoError(t, config.SetEnabled(ctx, false)) + assert.Empty(t, impl.activeRules) + assert.Empty(t, config.temporaryRules) + assert.NoError(t, restoreContextErr) + assert.False(t, config.enabled) + require.Len(t, impl.calls, 2) + assert.NoError(t, impl.calls[1].contextErr) +} + +func Test_defaultRouteInterfacesForIP(t *testing.T) { + t.Parallel() + + defaultRoutes := []routing.DefaultRoute{ + {NetInterface: "eth0", Family: netlink.FamilyV4}, + {NetInterface: "eth0", Family: netlink.FamilyV4}, + {NetInterface: "eth1", Family: netlink.FamilyV4}, + {NetInterface: "eth2", Family: netlink.FamilyV6}, + } + testCases := map[string]struct { + ip netip.Addr + interfaces []string + }{ + "ipv4": { + ip: netip.MustParseAddr("100.100.100.100"), + interfaces: []string{"eth0", "eth1"}, + }, + "ipv6": { + ip: netip.MustParseAddr("2001:db8::53"), + interfaces: []string{"eth2"}, + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + interfaces := defaultRouteInterfacesForIP(defaultRoutes, testCase.ip) + assert.Equal(t, testCase.interfaces, interfaces) + }) + } +} From 94ad398e11d1aa715abc41d9549ac60483449e2d Mon Sep 17 00:00:00 2001 From: Neil Date: Mon, 13 Jul 2026 09:55:21 -0600 Subject: [PATCH 2/3] feat(pia): native WireGuard support with port forwarding Private Internet Access previously only supported OpenVPN in gluetun because PIA registers a fresh WireGuard key per connection via its own API, which does not fit a static WireGuard config. This adds a provider-driven dynamic WireGuard path for PIA. At connection time, for VPN_SERVICE_PROVIDER=private internet access + VPN_TYPE=wireguard, gluetun now: - fetches PIA's live server list and selects a WireGuard server for the chosen region/name/hostname (honouring port-forwarding-only), - obtains an auth token, generates an ephemeral Curve25519 key pair, and registers the public key with the selected server (addKey) over TLS pinned to the server CN using the bundled PIA CA, - builds the WireGuard connection from the response (endpoint, peer key, interface address, DNS), re-registering on every reconnect, - performs the pre-tunnel token/server-list/addKey calls through scoped, temporary, mark-tagged firewall allowances so the killswitch is never opened to arbitrary destinations, resolving via a bootstrap dialer pinned to the physical default route. Port forwarding now works natively over WireGuard for PIA (gateway taken from the addKey server_vip). The earlier custom-provider workaround env VPN_PORT_FORWARDING_SERVER_NAME remains supported. Invalid, unspecified or non-IPv4 server_ip/server_vip/peer_ip values are rejected with errors rather than panicking, and the persisted port-forward token file is written and kept 0600. Config is just credentials + region, e.g.: VPN_SERVICE_PROVIDER=private internet access VPN_TYPE=wireguard SERVER_REGIONS=CA Vancouver VPN_PORT_FORWARDING=on VPN_PORT_FORWARDING_PROVIDER=private internet access Closes #3070. --- .../configuration/settings/portforward.go | 13 + .../settings/portforward_test.go | 17 + internal/configuration/settings/provider.go | 10 + .../configuration/settings/serverselection.go | 6 + internal/configuration/settings/vpn.go | 23 ++ internal/configuration/settings/vpn_test.go | 108 ++++++ internal/configuration/settings/wireguard.go | 13 +- .../settings/wireguardselection.go | 4 +- internal/models/server.go | 51 +-- internal/models/server_test.go | 16 + internal/models/wireguard.go | 12 + internal/portforward/service/settings.go | 10 +- internal/portforward/service/settings_test.go | 30 ++ internal/portforward/service/start.go | 9 +- .../privateinternetaccess/connection.go | 80 ++++ .../privateinternetaccess/connection_test.go | 125 +++++++ .../privateinternetaccess/httpclient.go | 28 ++ .../privateinternetaccess/portforward.go | 45 ++- .../privateinternetaccess/portforward_test.go | 186 ++++++++++ .../privateinternetaccess/provider.go | 22 +- .../privateinternetaccess/updater/api.go | 19 +- .../updater/hosttoserver.go | 10 +- .../privateinternetaccess/updater/servers.go | 15 +- .../updater/servers_test.go | 49 +++ .../updater/wireguard.go | 100 +++++ .../updater/wireguard_test.go | 160 ++++++++ .../privateinternetaccess/wireguard.go | 243 +++++++++++++ .../privateinternetaccess/wireguard_test.go | 344 ++++++++++++++++++ internal/provider/provider.go | 19 + internal/storage/formatting.go | 14 +- internal/storage/formatting_test.go | 29 ++ internal/vpn/bootstrap.go | 113 ++++++ internal/vpn/bootstrap_dial_linux.go | 32 ++ internal/vpn/bootstrap_dial_other.go | 18 + internal/vpn/bootstrap_test.go | 122 +++++++ internal/vpn/interfaces.go | 2 + internal/vpn/portforward.go | 1 + internal/vpn/run.go | 11 +- internal/vpn/tunnelup.go | 1 + internal/vpn/wireguard.go | 93 ++++- internal/vpn/wireguard_test.go | 38 ++ internal/wireguard/config.go | 22 +- internal/wireguard/config_test.go | 7 +- 43 files changed, 2171 insertions(+), 99 deletions(-) create mode 100644 internal/configuration/settings/vpn_test.go create mode 100644 internal/models/wireguard.go create mode 100644 internal/portforward/service/settings_test.go create mode 100644 internal/provider/privateinternetaccess/connection_test.go create mode 100644 internal/provider/privateinternetaccess/updater/servers_test.go create mode 100644 internal/provider/privateinternetaccess/updater/wireguard.go create mode 100644 internal/provider/privateinternetaccess/updater/wireguard_test.go create mode 100644 internal/provider/privateinternetaccess/wireguard.go create mode 100644 internal/provider/privateinternetaccess/wireguard_test.go create mode 100644 internal/storage/formatting_test.go create mode 100644 internal/vpn/bootstrap.go create mode 100644 internal/vpn/bootstrap_dial_linux.go create mode 100644 internal/vpn/bootstrap_dial_other.go create mode 100644 internal/vpn/bootstrap_test.go diff --git a/internal/configuration/settings/portforward.go b/internal/configuration/settings/portforward.go index 71850b758..86c63ca32 100644 --- a/internal/configuration/settings/portforward.go +++ b/internal/configuration/settings/portforward.go @@ -53,6 +53,10 @@ type PortForwarding struct { Username string `json:"username"` // Password is only used for Private Internet Access port forwarding. Password string `json:"password"` + // ServerName is only used for Private Internet Access port forwarding + // TLS certificate verification. It overrides the server name obtained + // from the VPN connection when set. + ServerName string `json:"server_name"` } func (p PortForwarding) Validate(vpnProvider string) (err error) { @@ -129,6 +133,7 @@ func (p *PortForwarding) Copy() (copied PortForwarding) { ListeningPorts: gosettings.CopySlice(p.ListeningPorts), Username: p.Username, Password: p.Password, + ServerName: p.ServerName, } } @@ -141,6 +146,7 @@ func (p *PortForwarding) OverrideWith(other PortForwarding) { p.ListeningPorts = gosettings.OverrideWithSlice(p.ListeningPorts, other.ListeningPorts) p.Username = gosettings.OverrideWithComparable(p.Username, other.Username) p.Password = gosettings.OverrideWithComparable(p.Password, other.Password) + p.ServerName = gosettings.OverrideWithComparable(p.ServerName, other.ServerName) } func (p *PortForwarding) setDefaults() { @@ -198,6 +204,10 @@ func (p PortForwarding) toLinesNode() (node *gotree.Node) { credentialsNode.Appendf("Password: %s", gosettings.ObfuscateKey(p.Password)) } + if p.ServerName != "" { + node.Appendf("Server name: %s", p.ServerName) + } + return node } @@ -237,6 +247,9 @@ func (p *PortForwarding) read(r *reader.Reader) (err error) { return err } + p.ServerName = r.String("VPN_PORT_FORWARDING_SERVER_NAME", + reader.ForceLowercase(false)) + usernameKeys := []string{"VPN_PORT_FORWARDING_USERNAME", "OPENVPN_USER", "USER"} for _, key := range usernameKeys { p.Username = r.String(key, reader.ForceLowercase(false)) diff --git a/internal/configuration/settings/portforward_test.go b/internal/configuration/settings/portforward_test.go index bb79b4466..786674ce0 100644 --- a/internal/configuration/settings/portforward_test.go +++ b/internal/configuration/settings/portforward_test.go @@ -3,6 +3,8 @@ package settings import ( "testing" + "github.com/qdm12/gosettings/reader" + "github.com/qdm12/gosettings/reader/sources/env" "github.com/stretchr/testify/assert" ) @@ -17,3 +19,18 @@ func Test_PortForwarding_String(t *testing.T) { assert.Empty(t, s) } + +func Test_PortForwarding_read_serverName(t *testing.T) { + t.Parallel() + + source := env.New(env.Settings{ + Environ: []string{"VPN_PORT_FORWARDING_SERVER_NAME=vancouver439"}, + }) + settingsReader := reader.New(reader.Settings{Sources: []reader.Source{source}}) + + var settings PortForwarding + err := settings.read(settingsReader) + + assert.NoError(t, err) + assert.Equal(t, "vancouver439", settings.ServerName) +} diff --git a/internal/configuration/settings/provider.go b/internal/configuration/settings/provider.go index ccc3acac3..7e55d8fce 100644 --- a/internal/configuration/settings/provider.go +++ b/internal/configuration/settings/provider.go @@ -1,6 +1,7 @@ package settings import ( + "errors" "fmt" "slices" "sort" @@ -49,6 +50,7 @@ func (p *Provider) validate(vpnType string, filterChoicesGetter FilterChoicesGet providers.Ivpn, providers.Mullvad, providers.Nordvpn, + providers.PrivateInternetAccess, providers.Protonvpn, providers.Surfshark, providers.Windscribe, @@ -68,6 +70,14 @@ func (p *Provider) validate(vpnType string, filterChoicesGetter FilterChoicesGet return fmt.Errorf("port forwarding: %w", err) } + customPIAPortForwarding := *p.PortForwarding.Enabled && + p.Name == providers.Custom && + *p.PortForwarding.Provider == providers.PrivateInternetAccess + serverNameMissing := p.PortForwarding.ServerName == "" && len(p.ServerSelection.Names) == 0 + if customPIAPortForwarding && serverNameMissing { + return errors.New("port forwarding: server name is empty: set VPN_PORT_FORWARDING_SERVER_NAME") + } + return nil } diff --git a/internal/configuration/settings/serverselection.go b/internal/configuration/settings/serverselection.go index c3e6a6d6c..a0e2bf14e 100644 --- a/internal/configuration/settings/serverselection.go +++ b/internal/configuration/settings/serverselection.go @@ -128,6 +128,12 @@ func getLocationFilterChoices(vpnServiceProvider string, filterChoices models.FilterChoices, err error, ) { filterChoices = filterChoicesGetter.GetFilterChoices(vpnServiceProvider) + if vpnServiceProvider == providers.PrivateInternetAccess && ss.VPN == vpn.Wireguard { + // PIA Wireguard region IDs and server names only exist in the live API, + // not in the embedded OpenVPN-only server list used for validation. + filterChoices.Regions = ss.Regions + filterChoices.Names = ss.Names + } if vpnServiceProvider == providers.Surfshark { // // Retro compatibility diff --git a/internal/configuration/settings/vpn.go b/internal/configuration/settings/vpn.go index 26fcdf20b..2617a6586 100644 --- a/internal/configuration/settings/vpn.go +++ b/internal/configuration/settings/vpn.go @@ -1,8 +1,10 @@ package settings import ( + "errors" "fmt" + "github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/constants/vpn" "github.com/qdm12/gosettings" "github.com/qdm12/gosettings/reader" @@ -57,6 +59,12 @@ func (v *VPN) Validate(filterChoicesGetter FilterChoicesGetter, ipv6Supported bo return fmt.Errorf("OpenVPN settings: %w", err) } case vpn.Wireguard: + if v.Provider.Name == providers.PrivateInternetAccess { + err = v.validatePIAWireguard() + if err != nil { + return fmt.Errorf("Private Internet Access settings: %w", err) + } + } const amneziawg = false err := v.Wireguard.validate(v.Provider.Name, ipv6Supported, amneziawg) if err != nil { @@ -72,6 +80,21 @@ func (v *VPN) Validate(filterChoicesGetter FilterChoicesGetter, ipv6Supported bo return nil } +func (v VPN) validatePIAWireguard() error { + switch { + case *v.OpenVPN.User == "": + return errors.New("username is empty: set OPENVPN_USER") + case *v.OpenVPN.Password == "": + return errors.New("password is empty: set OPENVPN_PASSWORD") + case len(v.Provider.ServerSelection.Regions) == 0 && + len(v.Provider.ServerSelection.Names) == 0 && + len(v.Provider.ServerSelection.Hostnames) == 0: + return errors.New("server selection is empty: set SERVER_REGIONS, SERVER_NAMES or SERVER_HOSTNAMES") + default: + return nil + } +} + func (v *VPN) Copy() (copied VPN) { return VPN{ Type: v.Type, diff --git a/internal/configuration/settings/vpn_test.go b/internal/configuration/settings/vpn_test.go new file mode 100644 index 000000000..8ab5aa9cd --- /dev/null +++ b/internal/configuration/settings/vpn_test.go @@ -0,0 +1,108 @@ +package settings + +import ( + "testing" + + "github.com/qdm12/gluetun/internal/constants/providers" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "github.com/stretchr/testify/assert" +) + +type testFilterChoicesGetter struct { + choices models.FilterChoices +} + +func (f testFilterChoicesGetter) GetFilterChoices(string) models.FilterChoices { + return f.choices +} + +type testWarner struct{} + +func (testWarner) Warn(string) {} + +func Test_VPN_validatePIAWireguard(t *testing.T) { + t.Parallel() + + const username = "user" + const password = "password" + const region = "CA Vancouver" + testCases := map[string]struct { + username string + password string + regions []string + names []string + wireguardKey string + expectedErrPart string + }{ + "valid": { + username: username, + password: password, + regions: []string{region}, + }, + "valid_region_id": { + username: username, + password: password, + regions: []string{"ca_vancouver"}, + }, + "valid_live_server_name": { + username: username, + password: password, + names: []string{"vancouver439"}, + }, + "username_missing": { + password: password, + regions: []string{region}, + expectedErrPart: "username is empty: set OPENVPN_USER", + }, + "password_missing": { + username: username, + regions: []string{region}, + expectedErrPart: "password is empty: set OPENVPN_PASSWORD", + }, + "server_selection_missing": { + username: username, + password: password, + expectedErrPart: "server selection is empty", + }, + "static_private_key_set": { + username: username, + password: password, + regions: []string{region}, + wireguardKey: "not-used", + expectedErrPart: "private key must not be set", + }, + } + + filterChoicesGetter := testFilterChoicesGetter{ + choices: models.FilterChoices{Regions: []string{region}}, + } + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + vpnSettings := VPN{ + Type: vpn.Wireguard, + Provider: Provider{ + Name: providers.PrivateInternetAccess, + ServerSelection: ServerSelection{ + VPN: vpn.Wireguard, + Regions: testCase.regions, + Names: testCase.names, + }, + }, + } + vpnSettings.setDefaults() + *vpnSettings.OpenVPN.User = testCase.username + *vpnSettings.OpenVPN.Password = testCase.password + *vpnSettings.Wireguard.PrivateKey = testCase.wireguardKey + + err := vpnSettings.Validate(filterChoicesGetter, false, testWarner{}) + + if testCase.expectedErrPart == "" { + assert.NoError(t, err) + } else { + assert.ErrorContains(t, err, testCase.expectedErrPart) + } + }) + } +} diff --git a/internal/configuration/settings/wireguard.go b/internal/configuration/settings/wireguard.go index d4b44a412..a4bfe9965 100644 --- a/internal/configuration/settings/wireguard.go +++ b/internal/configuration/settings/wireguard.go @@ -53,11 +53,18 @@ var regexpInterfaceName = regexp.MustCompile(`^[a-zA-Z0-9_]+$`) // Validate validates Wireguard settings. // It should only be ran if the VPN type chosen is Wireguard or AmneziaWg. func (w Wireguard) validate(vpnProvider string, ipv6Supported, amneziawg bool) (err error) { + dynamicPIAWireguard := vpnProvider == providers.PrivateInternetAccess && !amneziawg + if dynamicPIAWireguard && *w.PrivateKey != "" { + return errors.New("private key must not be set for Private Internet Access") + } + // Validate PrivateKey - if *w.PrivateKey == "" { + if *w.PrivateKey == "" && !dynamicPIAWireguard { return errors.New("private key is not set") } - _, err = wgtypes.ParseKey(*w.PrivateKey) + if *w.PrivateKey != "" { + _, err = wgtypes.ParseKey(*w.PrivateKey) + } if err != nil { err = fmt.Errorf("private key is not valid: %w", err) if vpnProvider == providers.Nordvpn && @@ -82,7 +89,7 @@ func (w Wireguard) validate(vpnProvider string, ipv6Supported, amneziawg bool) ( } // Validate Addresses - if len(w.Addresses) == 0 { + if len(w.Addresses) == 0 && !dynamicPIAWireguard { return errors.New("interface address is not set") } for i, ipNet := range w.Addresses { diff --git a/internal/configuration/settings/wireguardselection.go b/internal/configuration/settings/wireguardselection.go index 301893cf1..3f66b892d 100644 --- a/internal/configuration/settings/wireguardselection.go +++ b/internal/configuration/settings/wireguardselection.go @@ -41,7 +41,7 @@ func (w WireguardSelection) validate(vpnProvider string) (err error) { switch vpnProvider { case providers.Airvpn, providers.Fastestvpn, providers.Ivpn, providers.Mullvad, providers.Nordvpn, providers.Protonvpn, - providers.Surfshark, providers.Windscribe: + providers.PrivateInternetAccess, providers.Surfshark, providers.Windscribe: // endpoint IP addresses are baked in case providers.Custom: if !w.EndpointIP.IsValid() || w.EndpointIP.IsUnspecified() { @@ -58,7 +58,7 @@ func (w WireguardSelection) validate(vpnProvider string) (err error) { return errors.New("endpoint port is not set") } // EndpointPort cannot be set - case providers.Fastestvpn, providers.Nordvpn, + case providers.Fastestvpn, providers.Nordvpn, providers.PrivateInternetAccess, providers.Protonvpn, providers.Surfshark: if *w.EndpointPort != 0 { return errors.New("endpoint port is set") diff --git a/internal/models/server.go b/internal/models/server.go index 8c52ad663..c48dbbf4d 100644 --- a/internal/models/server.go +++ b/internal/models/server.go @@ -13,29 +13,30 @@ import ( type Server struct { VPN string `json:"vpn,omitempty"` // Surfshark: country is also used for multi-hop - Country string `json:"country,omitempty"` - Region string `json:"region,omitempty"` - City string `json:"city,omitempty"` - ISP string `json:"isp,omitempty"` - Categories []string `json:"categories,omitempty"` - Owned bool `json:"owned,omitempty"` - Number uint16 `json:"number,omitempty"` - ServerName string `json:"server_name,omitempty"` - Hostname string `json:"hostname,omitempty"` - TCP bool `json:"tcp,omitempty"` - UDP bool `json:"udp,omitempty"` - OvpnX509 string `json:"x509,omitempty"` - RetroLoc string `json:"retroloc,omitempty"` // TODO remove in v4 - MultiHop bool `json:"multihop,omitempty"` - WgPubKey string `json:"wgpubkey,omitempty"` - Free bool `json:"free,omitempty"` // TODO v4 create a SubscriptionTier struct - Premium bool `json:"premium,omitempty"` - Stream bool `json:"stream,omitempty"` // TODO v4 create a Features struct - SecureCore bool `json:"secure_core,omitempty"` - Tor bool `json:"tor,omitempty"` - PortForward bool `json:"port_forward,omitempty"` - Keep bool `json:"keep,omitempty"` - IPs []netip.Addr `json:"ips,omitempty"` + Country string `json:"country,omitempty"` + Region string `json:"region,omitempty"` + City string `json:"city,omitempty"` + ISP string `json:"isp,omitempty"` + Categories []string `json:"categories,omitempty"` + Owned bool `json:"owned,omitempty"` + Number uint16 `json:"number,omitempty"` + ServerName string `json:"server_name,omitempty"` + Hostname string `json:"hostname,omitempty"` + TCP bool `json:"tcp,omitempty"` + UDP bool `json:"udp,omitempty"` + OvpnX509 string `json:"x509,omitempty"` + RetroLoc string `json:"retroloc,omitempty"` // TODO remove in v4 + MultiHop bool `json:"multihop,omitempty"` + WgPubKey string `json:"wgpubkey,omitempty"` + WireguardDynamic bool `json:"wireguard_dynamic,omitempty"` + Free bool `json:"free,omitempty"` // TODO v4 create a SubscriptionTier struct + Premium bool `json:"premium,omitempty"` + Stream bool `json:"stream,omitempty"` // TODO v4 create a Features struct + SecureCore bool `json:"secure_core,omitempty"` + Tor bool `json:"tor,omitempty"` + PortForward bool `json:"port_forward,omitempty"` + Keep bool `json:"keep,omitempty"` + IPs []netip.Addr `json:"ips,omitempty"` } func (s *Server) HasMinimumInformation() (err error) { @@ -48,7 +49,9 @@ func (s *Server) HasMinimumInformation() (err error) { return errors.New("no network protocol should be set") case s.VPN == vpn.OpenVPN && !s.TCP && !s.UDP: return errors.New("both TCP and UDP fields are false for OpenVPN") - case s.VPN == vpn.Wireguard && s.WgPubKey == "": + case s.VPN == vpn.Wireguard && s.WireguardDynamic && s.ServerName == "": + return errors.New("server name field is empty for dynamic wireguard") + case s.VPN == vpn.Wireguard && s.WgPubKey == "" && !s.WireguardDynamic: return errors.New("wireguard public key field is empty") default: return nil diff --git a/internal/models/server_test.go b/internal/models/server_test.go index 8cdd7de16..d77d6e1bc 100644 --- a/internal/models/server_test.go +++ b/internal/models/server_test.go @@ -4,6 +4,7 @@ import ( "net/netip" "testing" + "github.com/qdm12/gluetun/internal/constants/vpn" "github.com/stretchr/testify/assert" ) @@ -123,3 +124,18 @@ func Test_Server_Equal(t *testing.T) { }) } } + +func Test_Server_HasMinimumInformation_dynamicWireguard(t *testing.T) { + t.Parallel() + + server := Server{ + VPN: vpn.Wireguard, + ServerName: "vancouver439", + WireguardDynamic: true, + IPs: []netip.Addr{netip.MustParseAddr("198.51.100.2")}, + } + + err := server.HasMinimumInformation() + + assert.NoError(t, err) +} diff --git a/internal/models/wireguard.go b/internal/models/wireguard.go new file mode 100644 index 000000000..7c855021a --- /dev/null +++ b/internal/models/wireguard.go @@ -0,0 +1,12 @@ +package models + +import "net/netip" + +// WireguardConnection contains provider-generated Wireguard connection data. +type WireguardConnection struct { + Connection Connection + PrivateKey string + Addresses []netip.Prefix + DNSServers []netip.Addr + Gateway netip.Addr +} diff --git a/internal/portforward/service/settings.go b/internal/portforward/service/settings.go index 61ded963a..e1c05abc3 100644 --- a/internal/portforward/service/settings.go +++ b/internal/portforward/service/settings.go @@ -3,6 +3,7 @@ package service import ( "errors" "fmt" + "net/netip" "slices" "github.com/qdm12/gluetun/internal/constants/providers" @@ -16,6 +17,7 @@ type Settings struct { UpCommand string DownCommand string Interface string // needed for PIA, PrivateVPN and ProtonVPN, tun0 for example + Gateway netip.Addr ServerName string // needed for PIA CanPortForward bool // needed for PIA ListeningPorts []uint16 @@ -31,6 +33,7 @@ func (s Settings) Copy() (copied Settings) { copied.UpCommand = s.UpCommand copied.DownCommand = s.DownCommand copied.Interface = s.Interface + copied.Gateway = s.Gateway copied.ServerName = s.ServerName copied.CanPortForward = s.CanPortForward copied.ListeningPorts = gosettings.CopySlice(s.ListeningPorts) @@ -47,8 +50,11 @@ func (s *Settings) OverrideWith(update Settings) { s.UpCommand = gosettings.OverrideWithComparable(s.UpCommand, update.UpCommand) s.DownCommand = gosettings.OverrideWithComparable(s.DownCommand, update.DownCommand) s.Interface = gosettings.OverrideWithComparable(s.Interface, update.Interface) - s.ServerName = gosettings.OverrideWithComparable(s.ServerName, update.ServerName) - s.CanPortForward = gosettings.OverrideWithComparable(s.CanPortForward, update.CanPortForward) + // These are runtime connection states, so zero values deliberately clear + // values from a previous connection. + s.Gateway = update.Gateway + s.ServerName = update.ServerName + s.CanPortForward = update.CanPortForward s.ListeningPorts = gosettings.OverrideWithSlice(s.ListeningPorts, update.ListeningPorts) s.PortsCount = gosettings.OverrideWithComparable(s.PortsCount, update.PortsCount) s.Username = gosettings.OverrideWithComparable(s.Username, update.Username) diff --git a/internal/portforward/service/settings_test.go b/internal/portforward/service/settings_test.go new file mode 100644 index 000000000..585820243 --- /dev/null +++ b/internal/portforward/service/settings_test.go @@ -0,0 +1,30 @@ +package service + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" +) + +func Test_Settings_OverrideWith_replacesConnectionRuntimeState(t *testing.T) { + t.Parallel() + + settings := Settings{ + Gateway: netip.MustParseAddr("10.0.0.1"), + ServerName: "old-server", + CanPortForward: true, + Username: "username", + Password: "password", + PortsCount: 1, + } + + settings.OverrideWith(Settings{}) + + assert.False(t, settings.Gateway.IsValid()) + assert.Empty(t, settings.ServerName) + assert.False(t, settings.CanPortForward) + assert.Equal(t, "username", settings.Username) + assert.Equal(t, "password", settings.Password) + assert.Equal(t, uint16(1), settings.PortsCount) +} diff --git a/internal/portforward/service/start.go b/internal/portforward/service/start.go index 7ee6875a4..8f1426e16 100644 --- a/internal/portforward/service/start.go +++ b/internal/portforward/service/start.go @@ -20,9 +20,12 @@ func (s *Service) Start(ctx context.Context) (runError <-chan error, err error) s.logger.Info("starting") - gateway, err := s.routing.VPNLocalGatewayIP(s.settings.Interface) - if err != nil { - return nil, fmt.Errorf("getting VPN local gateway IP: %w", err) + gateway := s.settings.Gateway + if !gateway.IsValid() { + gateway, err = s.routing.VPNLocalGatewayIP(s.settings.Interface) + if err != nil { + return nil, fmt.Errorf("getting VPN local gateway IP: %w", err) + } } family := netlink.FamilyV4 diff --git a/internal/provider/privateinternetaccess/connection.go b/internal/provider/privateinternetaccess/connection.go index 59dd866ec..8210966d7 100644 --- a/internal/provider/privateinternetaccess/connection.go +++ b/internal/provider/privateinternetaccess/connection.go @@ -1,17 +1,97 @@ package privateinternetaccess import ( + "context" + "errors" + "fmt" + "net" + "net/netip" + "time" + "github.com/qdm12/gluetun/internal/configuration/settings" + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/constants/vpn" "github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/provider/privateinternetaccess/presets" + "github.com/qdm12/gluetun/internal/provider/privateinternetaccess/updater" "github.com/qdm12/gluetun/internal/provider/utils" ) +// GetWireguardConnection obtains a live PIA Wireguard registration endpoint. +// The embedded PIA server list only contains OpenVPN servers. +func (p *Provider) GetWireguardConnection(ctx context.Context, selection settings.ServerSelection, + lookupNetIP func(context.Context, string, string) ([]netip.Addr, error), + dialContext func(context.Context, string, string) (net.Conn, error), + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error), +) ( + connection models.Connection, err error, +) { + const serverListHostname = "serverlist.piaservers.net" + const lookupTimeout = 10 * time.Second + lookupCtx, cancel := context.WithTimeout(ctx, lookupTimeout) + defer cancel() + addresses, err := lookupNetIP(lookupCtx, "ip4", serverListHostname) + if err != nil { + return connection, fmt.Errorf("resolving PIA server list host: %w", err) + } + if len(addresses) == 0 { + return connection, errors.New("resolving PIA server list host: no IPv4 address found") + } + + const serverListPort uint16 = 443 + serverListConnection := models.Connection{ + IP: addresses[0], + Port: serverListPort, + Protocol: constants.TCP, + } + removeConnection, err := allowConnection(ctx, serverListConnection) + if err != nil { + return connection, fmt.Errorf("allowing PIA server list connection: %w", err) + } + defer cleanupTemporaryConnection(ctx, removeConnection, &err) + + client, err := p.newDialingClient(serverListHostname, serverListConnection.IP, dialContext) + if err != nil { + return connection, fmt.Errorf("creating PIA server list client: %w", err) + } + defer client.CloseIdleConnections() + + liveUpdater := updater.New(client) + server, err := liveUpdater.FetchWireguardServer(ctx, selection) + if err != nil { + return connection, fmt.Errorf("fetching Wireguard server: %w", err) + } + + return models.Connection{ + Type: vpn.Wireguard, + IP: server.IPs[0], + Port: wireguardRegistrationPort, + Protocol: constants.UDP, + Hostname: server.Hostname, + ServerName: server.ServerName, + PortForward: server.PortForward, + }, nil +} + +func cleanupTemporaryConnection(ctx context.Context, remove func(context.Context) error, err *error) { + removeErr := remove(context.WithoutCancel(ctx)) + if removeErr != nil { + removeErr = fmt.Errorf("removing temporary connection allowance: %w", removeErr) + *err = errors.Join(*err, removeErr) + } +} + func (p *Provider) GetConnection(selection settings.ServerSelection, ipv6Supported bool) ( connection models.Connection, err error, ) { // Set port defaults depending on encryption preset. var defaults utils.ConnectionDefaults + if selection.VPN == vpn.Wireguard { + defaults.WireguardPort = wireguardRegistrationPort + return utils.GetConnection(p.Name(), + p.storage, selection, defaults, ipv6Supported, p.connPicker) + } + switch *selection.OpenVPN.PIAEncPreset { case presets.Normal: defaults.OpenVPNTCPPort = 502 diff --git a/internal/provider/privateinternetaccess/connection_test.go b/internal/provider/privateinternetaccess/connection_test.go new file mode 100644 index 000000000..0f157e6b1 --- /dev/null +++ b/internal/provider/privateinternetaccess/connection_test.go @@ -0,0 +1,125 @@ +package privateinternetaccess + +import ( + "context" + "io" + "net" + "net/http" + "net/netip" + "strings" + "testing" + "time" + + "github.com/golang/mock/gomock" + "github.com/qdm12/gluetun/internal/configuration/settings" + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/constants/providers" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "github.com/qdm12/gluetun/internal/provider/common" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_Provider_GetWireguardConnection(t *testing.T) { + t.Parallel() + + client := &http.Client{ + Transport: piaRoundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Body: io.NopCloser(strings.NewReader( + `{"regions":[{"id":"ca_vancouver","name":"CA Vancouver",` + + `"dns":"ca-vancouver.privacy.network","port_forward":true,` + + `"servers":{"wg":[{"ip":"198.51.100.2","cn":"vancouver439"}]}}]}` + + "\nsignature")), + }, nil + }), + } + provider := New(nil, time.Now, client) + const serverListIPString = "192.0.2.10" + serverListIP := netip.MustParseAddr(serverListIPString) + lookupNetIP := func(_ context.Context, network, host string) ([]netip.Addr, error) { + assert.Equal(t, "ip4", network) + assert.Equal(t, "serverlist.piaservers.net", host) + return []netip.Addr{serverListIP}, nil + } + provider.newDialingClient = func(serverName string, serverIP netip.Addr, + _ func(context.Context, string, string) (net.Conn, error), + ) (*http.Client, error) { + assert.Equal(t, "serverlist.piaservers.net", serverName) + assert.Equal(t, serverListIP, serverIP) + return client, nil + } + selection := settings.ServerSelection{ + VPN: vpn.Wireguard, + Regions: []string{"ca_vancouver"}, + }.WithDefaults(providers.PrivateInternetAccess) + + var allowedConnection models.Connection + connectionAllowanceRemoved := false + allowConnection := func(_ context.Context, connection models.Connection) ( + func(context.Context) error, error, + ) { + allowedConnection = connection + return func(context.Context) error { + connectionAllowanceRemoved = true + return nil + }, nil + } + connection, err := provider.GetWireguardConnection(context.Background(), selection, + lookupNetIP, new(net.Dialer).DialContext, allowConnection) + require.NoError(t, err) + assert.Equal(t, models.Connection{ + IP: serverListIP, + Port: 443, + Protocol: constants.TCP, + }, allowedConnection) + assert.Equal(t, models.Connection{ + Type: vpn.Wireguard, + IP: netip.MustParseAddr("198.51.100.2"), + Port: wireguardRegistrationPort, + Protocol: constants.UDP, + Hostname: "ca-vancouver.privacy.network", + ServerName: "vancouver439", + PortForward: true, + }, connection) + assert.True(t, connectionAllowanceRemoved) +} + +func Test_Provider_GetConnection_wireguard(t *testing.T) { + t.Parallel() + + selection := settings.ServerSelection{VPN: vpn.Wireguard}. + WithDefaults(providers.PrivateInternetAccess) + server := models.Server{ + VPN: vpn.Wireguard, + Region: "CA Vancouver", + ServerName: testPIAServerName, + Hostname: "ca-vancouver.privacy.network", + PortForward: true, + WireguardDynamic: true, + IPs: []netip.Addr{netip.MustParseAddr("198.51.100.2")}, + } + controller := gomock.NewController(t) + storage := common.NewMockStorage(controller) + storage.EXPECT().FilterServers(providers.PrivateInternetAccess, selection). + Return([]models.Server{server}, nil) + provider := New(storage, time.Now, nil) + + connection, err := provider.GetConnection(selection, false) + require.NoError(t, err) + + const registrationPort uint16 = 1337 + assert.Equal(t, registrationPort, connection.Port) + assert.Equal(t, server.IPs[0], connection.IP) + assert.Equal(t, server.ServerName, connection.ServerName) + assert.True(t, connection.PortForward) +} + +type piaRoundTripFunc func(request *http.Request) (*http.Response, error) + +func (f piaRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return f(request) +} diff --git a/internal/provider/privateinternetaccess/httpclient.go b/internal/provider/privateinternetaccess/httpclient.go index ebd076791..7eda77369 100644 --- a/internal/provider/privateinternetaccess/httpclient.go +++ b/internal/provider/privateinternetaccess/httpclient.go @@ -1,11 +1,13 @@ package privateinternetaccess import ( + "context" "crypto/tls" "crypto/x509" "fmt" "net" "net/http" + "net/netip" "strings" "time" @@ -48,3 +50,29 @@ func newHTTPClient(serverName string) (client *http.Client, err error) { Timeout: 30 * time.Second, }, nil } + +func newHTTPClientDialing(serverName string, serverIP netip.Addr, + dialContext func(context.Context, string, string) (net.Conn, error), +) (client *http.Client, err error) { + client, err = newHTTPClient(serverName) + if err != nil { + return nil, err + } + + transport, ok := client.Transport.(*http.Transport) + if !ok { + panic("PIA HTTP client transport has an unexpected type") + } + + transport.Proxy = nil + transport.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + _, port, err := net.SplitHostPort(address) + if err != nil { + return nil, fmt.Errorf("splitting dial address: %w", err) + } + address = net.JoinHostPort(serverIP.String(), port) + return dialContext(ctx, network, address) + } + + return client, nil +} diff --git a/internal/provider/privateinternetaccess/portforward.go b/internal/provider/privateinternetaccess/portforward.go index 9c8d9abb9..267f5a17f 100644 --- a/internal/provider/privateinternetaccess/portforward.go +++ b/internal/provider/privateinternetaccess/portforward.go @@ -28,13 +28,13 @@ func (p *Provider) PortForward(ctx context.Context, ) (internalToExternalPorts map[uint16]uint16, err error) { switch { case objects.ServerName == "": - panic("server name cannot be empty") + return nil, errors.New("server name cannot be empty") case !objects.Gateway.IsValid(): - panic("gateway is not set") + return nil, errors.New("gateway is not set") case objects.Username == "": - panic("username is not set") + return nil, errors.New("username is not set") case objects.Password == "": - panic("password is not set") + return nil, errors.New("password is not set") } serverName := objects.ServerName @@ -109,9 +109,9 @@ func (p *Provider) KeepPortForward(ctx context.Context, ) (err error) { switch { case objects.ServerName == "": - panic("server name cannot be empty") + return errors.New("server name cannot be empty") case !objects.Gateway.IsValid(): - panic("gateway is not set") + return errors.New("gateway is not set") } privateIPClient, err := newHTTPClient(objects.ServerName) @@ -156,8 +156,9 @@ func (p *Provider) KeepPortForward(ctx context.Context, func findAPIIP(ctx context.Context, client *http.Client, gateway netip.Addr) ( apiIP netip.Addr, err error, ) { - if gateway.Is6() { - panic("IPv6 gateway not supported") + err = validateRegistrationIPv4(gateway, "gateway") + if err != nil { + return netip.Addr{}, err } gatewayBytes := gateway.As4() @@ -240,6 +241,12 @@ func readPIAPortForwardData(portForwardPath string) (data piaPortForwardData, er } else if err != nil { return data, err } + const permission = fs.FileMode(0o600) + err = file.Chmod(permission) + if err != nil { + _ = file.Close() + return data, err + } decoder := json.NewDecoder(file) if err := decoder.Decode(&data); err != nil { @@ -251,11 +258,16 @@ func readPIAPortForwardData(portForwardPath string) (data piaPortForwardData, er } func writePIAPortForwardData(portForwardPath string, data piaPortForwardData) (err error) { - const permission = fs.FileMode(0o644) + const permission = fs.FileMode(0o600) file, err := os.OpenFile(portForwardPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, permission) if err != nil { return err } + err = file.Chmod(permission) + if err != nil { + _ = file.Close() + return err + } encoder := json.NewEncoder(file) @@ -299,8 +311,16 @@ func packPayload(port uint16, token string, expiration time.Time) (payload strin return payload, nil } +const piaTokenURL = "https://www.privateinternetaccess.com/api/client/v2/token" + func fetchToken(ctx context.Context, client *http.Client, username, password string, +) (token string, err error) { + return fetchTokenFromURL(ctx, client, piaTokenURL, username, password) +} + +func fetchTokenFromURL(ctx context.Context, client *http.Client, + tokenURL, username, password string, ) (token string, err error) { errSubstitutions := map[string]string{ url.QueryEscape(username): "", @@ -316,12 +336,7 @@ func fetchToken(ctx context.Context, client *http.Client, form := url.Values{} form.Add("username", username) form.Add("password", password) - url := url.URL{ - Scheme: "https", - Host: "www.privateinternetaccess.com", - Path: "/api/client/v2/token", - } - request, err := http.NewRequestWithContext(ctx, http.MethodPost, url.String(), strings.NewReader(form.Encode())) + request, err := http.NewRequestWithContext(ctx, http.MethodPost, tokenURL, strings.NewReader(form.Encode())) if err != nil { return "", replaceInErr(err, errSubstitutions) } diff --git a/internal/provider/privateinternetaccess/portforward_test.go b/internal/provider/privateinternetaccess/portforward_test.go index 371c92693..c56681ef3 100644 --- a/internal/provider/privateinternetaccess/portforward_test.go +++ b/internal/provider/privateinternetaccess/portforward_test.go @@ -1,16 +1,202 @@ package privateinternetaccess import ( + "context" "encoding/base64" "encoding/json" "errors" + "net/http" + "net/http/httptest" + "net/netip" + "os" "testing" "time" + "github.com/qdm12/gluetun/internal/provider/utils" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) +const testPIAServerName = "vancouver439" + +func Test_fetchToken(t *testing.T) { + t.Parallel() + + type receivedRequest struct { + method string + username string + password string + } + received := make(chan receivedRequest, 1) + server := httptest.NewServer(http.HandlerFunc(func(responseWriter http.ResponseWriter, request *http.Request) { + _ = request.ParseForm() + received <- receivedRequest{ + method: request.Method, + username: request.Form.Get("username"), + password: request.Form.Get("password"), + } + responseWriter.Header().Set("Content-Type", "application/json") + _, _ = responseWriter.Write([]byte(`{"token":"test-token"}`)) + })) + t.Cleanup(server.Close) + + token, err := fetchTokenFromURL(context.Background(), server.Client(), + server.URL, "test-user", "test-password") + require.NoError(t, err) + + request := <-received + assert.Equal(t, http.MethodPost, request.method) + assert.Equal(t, "test-user", request.username) + assert.Equal(t, "test-password", request.password) + assert.Equal(t, "test-token", token) +} + +func Test_PortForward_inputValidation(t *testing.T) { + t.Parallel() + + testCases := map[string]struct { + objects utils.PortForwardObjects + errMessage string + }{ + "server_name_not_set": { + errMessage: "server name cannot be empty", + }, + "gateway_not_set": { + objects: utils.PortForwardObjects{ + ServerName: testPIAServerName, + }, + errMessage: "gateway is not set", + }, + "username_not_set": { + objects: utils.PortForwardObjects{ + ServerName: testPIAServerName, + Gateway: netip.MustParseAddr("10.13.161.1"), + }, + errMessage: "username is not set", + }, + "password_not_set": { + objects: utils.PortForwardObjects{ + ServerName: testPIAServerName, + Gateway: netip.MustParseAddr("10.13.161.1"), + Username: "username", + }, + errMessage: "password is not set", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + provider := &Provider{} + + _, err := provider.PortForward(context.Background(), testCase.objects) + + assert.ErrorContains(t, err, testCase.errMessage) + }) + } +} + +func Test_KeepPortForward_inputValidation(t *testing.T) { + t.Parallel() + + testCases := map[string]struct { + objects utils.PortForwardObjects + errMessage string + }{ + "server_name_not_set": { + errMessage: "server name cannot be empty", + }, + "gateway_not_set": { + objects: utils.PortForwardObjects{ + ServerName: testPIAServerName, + }, + errMessage: "gateway is not set", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + provider := &Provider{} + + err := provider.KeepPortForward(context.Background(), testCase.objects) + + assert.ErrorContains(t, err, testCase.errMessage) + }) + } +} + +func Test_findAPIIP_rejectsUnsupportedGateway(t *testing.T) { + t.Parallel() + + testCases := map[string]struct { + gateway netip.Addr + errMessage string + }{ + "not_set": { + errMessage: "gateway is not set", + }, + "unspecified": { + gateway: netip.IPv4Unspecified(), + errMessage: "gateway is unspecified", + }, + "ipv6": { + gateway: netip.MustParseAddr("2001:db8::1"), + errMessage: "gateway is IPv6, which PIA registration does not support", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + _, err := findAPIIP(t.Context(), nil, testCase.gateway) + + require.Error(t, err) + assert.ErrorContains(t, err, testCase.errMessage) + }) + } +} + +func Test_readPIAPortForwardData_restrictsPermissionsOnReuse(t *testing.T) { + t.Parallel() + + expectedData := piaPortForwardData{ + Port: 12345, + Token: "secret-token", + Signature: "signature", + Expiration: time.Now().Add(time.Hour).UTC().Truncate(time.Second), + } + contents, err := json.Marshal(expectedData) + require.NoError(t, err) + path := t.TempDir() + "/pia.json" + require.NoError(t, os.WriteFile(path, contents, 0o644)) + require.NoError(t, os.Chmod(path, 0o644)) + + data, err := readPIAPortForwardData(path) + require.NoError(t, err) + assert.Equal(t, expectedData, data) + + fileInfo, err := os.Stat(path) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0o600), fileInfo.Mode().Perm()) +} + +func Test_writePIAPortForwardData_restrictsPermissions(t *testing.T) { + t.Parallel() + + path := t.TempDir() + "/pia.json" + require.NoError(t, os.WriteFile(path, []byte("old data"), 0o644)) + require.NoError(t, os.Chmod(path, 0o644)) + + err := writePIAPortForwardData(path, piaPortForwardData{Token: "secret"}) + require.NoError(t, err) + + fileInfo, err := os.Stat(path) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0o600), fileInfo.Mode().Perm()) +} + func Test_unpackPayload(t *testing.T) { t.Parallel() diff --git a/internal/provider/privateinternetaccess/provider.go b/internal/provider/privateinternetaccess/provider.go index da5ecf783..2cddd2978 100644 --- a/internal/provider/privateinternetaccess/provider.go +++ b/internal/provider/privateinternetaccess/provider.go @@ -1,6 +1,8 @@ package privateinternetaccess import ( + "context" + "net" "net/http" "net/netip" "time" @@ -12,9 +14,11 @@ import ( ) type Provider struct { - storage common.Storage - connPicker *utils.ConnectionPicker - timeNow func() time.Time + storage common.Storage + connPicker *utils.ConnectionPicker + timeNow func() time.Time + newDialingClient func(string, netip.Addr, + func(context.Context, string, string) (net.Conn, error)) (*http.Client, error) common.Fetcher // Port forwarding portForwardPath string @@ -25,12 +29,14 @@ func New(storage common.Storage, timeNow func() time.Time, client *http.Client, ) *Provider { const jsonPortForwardPath = "/gluetun/piaportforward.json" + serverUpdater := updater.New(client) return &Provider{ - storage: storage, - timeNow: timeNow, - connPicker: utils.NewConnectionPicker(), - portForwardPath: jsonPortForwardPath, - Fetcher: updater.New(client), + storage: storage, + timeNow: timeNow, + connPicker: utils.NewConnectionPicker(), + newDialingClient: newHTTPClientDialing, + portForwardPath: jsonPortForwardPath, + Fetcher: serverUpdater, } } diff --git a/internal/provider/privateinternetaccess/updater/api.go b/internal/provider/privateinternetaccess/updater/api.go index 584a470fe..0a79fcd67 100644 --- a/internal/provider/privateinternetaccess/updater/api.go +++ b/internal/provider/privateinternetaccess/updater/api.go @@ -15,6 +15,7 @@ type apiData struct { } type regionData struct { + ID string `json:"id"` Name string `json:"name"` DNS string `json:"dns"` PortForward bool `json:"port_forward"` @@ -22,6 +23,7 @@ type regionData struct { Servers struct { UDP []serverData `json:"ovpnudp"` TCP []serverData `json:"ovpntcp"` + WG []serverData `json:"wg"` } `json:"servers"` } @@ -34,7 +36,19 @@ func fetchAPI(ctx context.Context, client *http.Client) ( data apiData, err error, ) { const url = "https://serverlist.piaservers.net/vpninfo/servers/v7" + return fetchAPIFromURL(ctx, client, url) +} +func fetchWireguardAPI(ctx context.Context, client *http.Client) ( + data apiData, err error, +) { + const url = "https://serverlist.piaservers.net/vpninfo/servers/v6" + return fetchAPIFromURL(ctx, client, url) +} + +func fetchAPIFromURL(ctx context.Context, client *http.Client, url string) ( + data apiData, err error, +) { request, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) if err != nil { return data, err @@ -60,9 +74,8 @@ func fetchAPI(ctx context.Context, client *http.Client) ( return data, err } - // remove key/signature at the bottom - i := bytes.IndexRune(b, '\n') - b = b[:i] + // Remove the key/signature after the JSON first line. + b, _, _ = bytes.Cut(b, []byte{'\n'}) if err := json.Unmarshal(b, &data); err != nil { return data, err diff --git a/internal/provider/privateinternetaccess/updater/hosttoserver.go b/internal/provider/privateinternetaccess/updater/hosttoserver.go index 93464106f..053065053 100644 --- a/internal/provider/privateinternetaccess/updater/hosttoserver.go +++ b/internal/provider/privateinternetaccess/updater/hosttoserver.go @@ -9,17 +9,19 @@ import ( type nameToServer map[string]models.Server -func (nts nameToServer) add(name, hostname, region string, +func (nts nameToServer) add(vpnType, name, hostname, region string, tcp, udp, portForward bool, ip netip.Addr, ) (change bool) { - server, ok := nts[name] + key := vpnType + "-" + name + server, ok := nts[key] if !ok { change = true - server.VPN = vpn.OpenVPN + server.VPN = vpnType server.ServerName = name server.Hostname = hostname server.Region = region server.PortForward = portForward + server.WireguardDynamic = vpnType == vpn.Wireguard } if !server.TCP && tcp { @@ -44,7 +46,7 @@ func (nts nameToServer) add(name, hostname, region string, server.IPs = append(server.IPs, ip) } - nts[name] = server + nts[key] = server return change } diff --git a/internal/provider/privateinternetaccess/updater/servers.go b/internal/provider/privateinternetaccess/updater/servers.go index c2b813044..124614c48 100644 --- a/internal/provider/privateinternetaccess/updater/servers.go +++ b/internal/provider/privateinternetaccess/updater/servers.go @@ -6,6 +6,7 @@ import ( "sort" "time" + "github.com/qdm12/gluetun/internal/constants/vpn" "github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/provider/common" ) @@ -82,14 +83,24 @@ func addData(regions []regionData, nts nameToServer) (change bool) { } for _, server := range region.Servers.UDP { const tcp, udp = false, true - if nts.add(server.CN, region.DNS, region.Name, tcp, udp, region.PortForward, server.IP) { + if nts.add(vpn.OpenVPN, server.CN, region.DNS, region.Name, + tcp, udp, region.PortForward, server.IP) { change = true } } for _, server := range region.Servers.TCP { const tcp, udp = true, false - if nts.add(server.CN, region.DNS, region.Name, tcp, udp, region.PortForward, server.IP) { + if nts.add(vpn.OpenVPN, server.CN, region.DNS, region.Name, + tcp, udp, region.PortForward, server.IP) { + change = true + } + } + + for _, server := range region.Servers.WG { + const tcp, udp = false, false + if nts.add(vpn.Wireguard, server.CN, region.DNS, region.Name, + tcp, udp, region.PortForward, server.IP) { change = true } } diff --git a/internal/provider/privateinternetaccess/updater/servers_test.go b/internal/provider/privateinternetaccess/updater/servers_test.go new file mode 100644 index 000000000..7d8113c21 --- /dev/null +++ b/internal/provider/privateinternetaccess/updater/servers_test.go @@ -0,0 +1,49 @@ +package updater + +import ( + "net/netip" + "testing" + + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_addData_wireguard(t *testing.T) { + t.Parallel() + + region := regionData{ + Name: "CA Vancouver", + DNS: "ca-vancouver.privacy.network", + PortForward: true, + } + region.Servers.UDP = []serverData{{ + IP: netip.MustParseAddr("198.51.100.1"), + CN: "vancouver-openvpn", + }} + region.Servers.WG = []serverData{{ + IP: netip.MustParseAddr("198.51.100.2"), + CN: "vancouver439", + }} + serversByName := make(nameToServer) + + changed := addData([]regionData{region}, serversByName) + require.True(t, changed) + servers := serversByName.toServersSlice() + require.Len(t, servers, 2) + + serverByVPN := make(map[string]models.Server, len(servers)) + for _, server := range servers { + serverByVPN[server.VPN] = server + } + wireguardServer := serverByVPN[vpn.Wireguard] + assert.Equal(t, "vancouver439", wireguardServer.ServerName) + assert.Equal(t, region.Name, wireguardServer.Region) + assert.Equal(t, region.DNS, wireguardServer.Hostname) + assert.Equal(t, []netip.Addr{netip.MustParseAddr("198.51.100.2")}, wireguardServer.IPs) + assert.True(t, wireguardServer.PortForward) + assert.True(t, wireguardServer.WireguardDynamic) + assert.False(t, wireguardServer.TCP) + assert.False(t, wireguardServer.UDP) +} diff --git a/internal/provider/privateinternetaccess/updater/wireguard.go b/internal/provider/privateinternetaccess/updater/wireguard.go new file mode 100644 index 000000000..aada68b7a --- /dev/null +++ b/internal/provider/privateinternetaccess/updater/wireguard.go @@ -0,0 +1,100 @@ +package updater + +import ( + "context" + "errors" + "fmt" + "net/netip" + "strings" + + "github.com/qdm12/gluetun/internal/configuration/settings" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" +) + +// FetchWireguardServer obtains a Wireguard server directly from PIA's live +// server list. PIA Wireguard servers are registered dynamically and are not +// present in the embedded server list used by the other connection paths. +func (u *Updater) FetchWireguardServer(ctx context.Context, selection settings.ServerSelection) ( + server models.Server, err error, +) { + data, err := fetchWireguardAPI(ctx, u.client) + if err != nil { + return server, fmt.Errorf("fetching PIA server list: %w", err) + } + + return selectWireguardServer(data.Regions, selection) +} + +func selectWireguardServer(regions []regionData, selection settings.ServerSelection) ( + server models.Server, err error, +) { + for _, region := range regions { + if region.Offline || (*selection.PortForwardOnly && !region.PortForward) || + !matchesAnyRegion(region, selection.Regions) || + !matchesAnyHostname(region, selection.Hostnames) { + continue + } + + for _, wireguardServer := range region.Servers.WG { + if !matchesAnyServer(region, wireguardServer, selection.Names) { + continue + } + + return models.Server{ + VPN: vpn.Wireguard, + Region: region.Name, + ServerName: wireguardServer.CN, + Hostname: region.DNS, + WireguardDynamic: true, + PortForward: region.PortForward, + IPs: []netip.Addr{wireguardServer.IP}, + }, nil + } + } + + return server, errors.New("no Wireguard server found matching selection") +} + +func matchesAnyRegion(region regionData, regions []string) bool { + if len(regions) == 0 { + return true + } + + for _, selectedRegion := range regions { + if strings.EqualFold(region.Name, selectedRegion) || strings.EqualFold(region.ID, selectedRegion) { + return true + } + } + + return false +} + +func matchesAnyHostname(region regionData, hostnames []string) bool { + if len(hostnames) == 0 { + return true + } + + for _, selectedHostname := range hostnames { + if strings.EqualFold(region.DNS, selectedHostname) { + return true + } + } + + return false +} + +func matchesAnyServer(region regionData, server serverData, names []string) bool { + if len(names) == 0 { + return true + } + + for _, selectedName := range names { + if strings.EqualFold(server.CN, selectedName) || + strings.EqualFold(region.Name, selectedName) || strings.EqualFold(region.ID, selectedName) { + return true + } + } + + return false +} diff --git a/internal/provider/privateinternetaccess/updater/wireguard_test.go b/internal/provider/privateinternetaccess/updater/wireguard_test.go new file mode 100644 index 000000000..7948c7078 --- /dev/null +++ b/internal/provider/privateinternetaccess/updater/wireguard_test.go @@ -0,0 +1,160 @@ +package updater + +import ( + "context" + "io" + "net/http" + "net/netip" + "strings" + "testing" + + "github.com/qdm12/gluetun/internal/configuration/settings" + "github.com/qdm12/gluetun/internal/constants/providers" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_Updater_FetchWireguardServer(t *testing.T) { + t.Parallel() + + client := &http.Client{ + Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) { + assert.Equal(t, http.MethodGet, request.Method) + assert.Equal(t, "https://serverlist.piaservers.net/vpninfo/servers/v6", request.URL.String()) + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Body: io.NopCloser(strings.NewReader( + `{"regions":[{"id":"ca_vancouver","name":"CA Vancouver",` + + `"dns":"ca-vancouver.privacy.network","port_forward":true,` + + `"servers":{"wg":[{"ip":"198.51.100.2","cn":"vancouver439"}]}}]}` + + "\nserver-list-signature")), + }, nil + }), + } + updater := New(client) + selection := settings.ServerSelection{ + VPN: vpn.Wireguard, + Regions: []string{"ca vAnCoUvEr"}, + }.WithDefaults(providers.PrivateInternetAccess) + + server, err := updater.FetchWireguardServer(context.Background(), selection) + require.NoError(t, err) + assert.Equal(t, models.Server{ + VPN: vpn.Wireguard, + Region: "CA Vancouver", + ServerName: "vancouver439", + Hostname: "ca-vancouver.privacy.network", + WireguardDynamic: true, + PortForward: true, + IPs: []netip.Addr{netip.MustParseAddr("198.51.100.2")}, + }, server) +} + +func Test_selectWireguardServer(t *testing.T) { + t.Parallel() + + vancouverRegion := regionData{ + ID: "ca_vancouver", + Name: "CA Vancouver", + DNS: "ca-vancouver.privacy.network", + PortForward: true, + } + vancouverRegion.Servers.WG = []serverData{ + {IP: netip.MustParseAddr("198.51.100.2"), CN: "vancouver439"}, + {IP: netip.MustParseAddr("198.51.100.3"), CN: "vancouver440"}, + } + londonRegion := regionData{ + ID: "uk_london", + Name: "UK London", + DNS: "uk-london.privacy.network", + } + londonRegion.Servers.WG = []serverData{{ + IP: netip.MustParseAddr("203.0.113.2"), CN: "london401", + }} + regions := []regionData{vancouverRegion, londonRegion} + + testCases := map[string]struct { + selection settings.ServerSelection + portForwardOnly bool + expectedServer string + errMessage string + }{ + "region_name_case_insensitive": { + selection: settings.ServerSelection{Regions: []string{"ca vAnCoUvEr"}}, + expectedServer: "vancouver439", + }, + "region_id": { + selection: settings.ServerSelection{Regions: []string{"CA_VANCOUVER"}}, + expectedServer: "vancouver439", + }, + "server_name": { + selection: settings.ServerSelection{Names: []string{"VANCOUVER440"}}, + expectedServer: "vancouver440", + }, + "hostname_only": { + selection: settings.ServerSelection{ + Hostnames: []string{"UK-LONDON.PRIVACY.NETWORK"}, + }, + expectedServer: "london401", + }, + "combined_region_name_and_hostname": { + selection: settings.ServerSelection{ + Regions: []string{"ca_vancouver"}, + Names: []string{"vancouver440"}, + Hostnames: []string{"ca-vancouver.privacy.network"}, + }, + expectedServer: "vancouver440", + }, + "hostname_no_match": { + selection: settings.ServerSelection{ + Hostnames: []string{"sydney.privacy.network"}, + }, + errMessage: "no Wireguard server found matching selection", + }, + "combined_filters_must_all_match": { + selection: settings.ServerSelection{ + Regions: []string{"ca_vancouver"}, + Hostnames: []string{"uk-london.privacy.network"}, + }, + errMessage: "no Wireguard server found matching selection", + }, + "region_id_as_server_name": { + selection: settings.ServerSelection{Names: []string{"ca_vancouver"}}, + expectedServer: "vancouver439", + }, + "port_forwarding_only": { + selection: settings.ServerSelection{Regions: []string{"UK London"}}, + portForwardOnly: true, + errMessage: "no Wireguard server found matching selection", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + selection := testCase.selection.WithDefaults(providers.PrivateInternetAccess) + if testCase.portForwardOnly { + *selection.PortForwardOnly = true + } + server, err := selectWireguardServer(regions, selection) + if testCase.errMessage != "" { + require.Error(t, err) + assert.ErrorContains(t, err, testCase.errMessage) + return + } + + require.NoError(t, err) + assert.Equal(t, testCase.expectedServer, server.ServerName) + } + } +} + +type roundTripFunc func(request *http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return f(request) +} diff --git a/internal/provider/privateinternetaccess/wireguard.go b/internal/provider/privateinternetaccess/wireguard.go new file mode 100644 index 000000000..40da57a8f --- /dev/null +++ b/internal/provider/privateinternetaccess/wireguard.go @@ -0,0 +1,243 @@ +package privateinternetaccess + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/http" + "net/netip" + "net/url" + "strconv" + "time" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +type addKeyResponse struct { + Status string `json:"status"` + ServerKey string `json:"server_key"` + ServerPort uint16 `json:"server_port"` + ServerIP netip.Addr `json:"server_ip"` + ServerVIP netip.Addr `json:"server_vip"` + PeerIP netip.Addr `json:"peer_ip"` + PeerPubKey string `json:"peer_pubkey"` + DNSServers []netip.Addr `json:"dns_servers"` +} + +const wireguardRegistrationPort uint16 = 1337 + +// RegisterWireguard registers a fresh ephemeral key with the selected PIA +// server and returns all provider-generated connection settings. +func (p *Provider) RegisterWireguard(ctx context.Context, connection models.Connection, + username, password string, + lookupNetIP func(context.Context, string, string) ([]netip.Addr, error), + dialContext func(context.Context, string, string) (net.Conn, error), + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error), +) (wireguardConnection models.WireguardConnection, err error) { + switch { + case !connection.IP.IsValid(): + return wireguardConnection, errors.New("registration server IP is not set") + case connection.ServerName == "": + return wireguardConnection, errors.New("registration server name is not set") + case username == "": + return wireguardConnection, errors.New("username is not set") + case password == "": + return wireguardConnection, errors.New("password is not set") + } + + token, err := p.fetchWireguardToken(ctx, username, password, + lookupNetIP, dialContext, allowConnection) + if err != nil { + return wireguardConnection, fmt.Errorf("fetching token: %w", err) + } + + privateKey, err := wgtypes.GeneratePrivateKey() + if err != nil { + return wireguardConnection, fmt.Errorf("generating Wireguard private key: %w", err) + } + publicKey := privateKey.PublicKey().String() + + client, err := p.newDialingClient(connection.ServerName, connection.IP, dialContext) + if err != nil { + return wireguardConnection, fmt.Errorf("creating registration HTTP client: %w", err) + } + + registrationConnection := connection + registrationConnection.Port = wireguardRegistrationPort + registrationConnection.Protocol = constants.TCP + removeConnection, err := allowConnection(ctx, registrationConnection) + if err != nil { + return wireguardConnection, fmt.Errorf("allowing registration connection: %w", err) + } + defer cleanupTemporaryConnection(ctx, removeConnection, &err) + defer client.CloseIdleConnections() + + response, err := fetchAddKey(ctx, client, connection.ServerName, + wireguardRegistrationPort, token, publicKey) + if err != nil { + return wireguardConnection, fmt.Errorf("registering Wireguard key: %w", err) + } + + wireguardConnection, err = mapAddKeyResponse(connection, privateKey.String(), response) + if err != nil { + return models.WireguardConnection{}, fmt.Errorf("mapping registration response: %w", err) + } + return wireguardConnection, nil +} + +func (p *Provider) fetchWireguardToken(ctx context.Context, username, password string, + lookupNetIP func(context.Context, string, string) ([]netip.Addr, error), + dialContext func(context.Context, string, string) (net.Conn, error), + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error), +) (token string, err error) { + const tokenServerName = "www.privateinternetaccess.com" + const timeout = 10 * time.Second + lookupCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + addresses, err := lookupNetIP(lookupCtx, "ip4", tokenServerName) + if err != nil { + return "", fmt.Errorf("resolving token server: %w", err) + } else if len(addresses) == 0 { + return "", errors.New("resolving token server: no IPv4 address found") + } + + const tokenServerPort = 443 + errs := make([]error, 0, len(addresses)) + for _, address := range addresses { + tokenConnection := models.Connection{ + IP: address, + Port: tokenServerPort, + Protocol: constants.TCP, + } + removeConnection, err := allowConnection(ctx, tokenConnection) + if err != nil { + errs = append(errs, fmt.Errorf("allowing token connection: %w", err)) + continue + } + + client, err := p.newDialingClient(tokenServerName, tokenConnection.IP, dialContext) + if err != nil { + cleanupTemporaryConnection(ctx, removeConnection, &err) + errs = append(errs, fmt.Errorf("creating token HTTP client: %w", err)) + continue + } + token, err = fetchTokenFromURL(ctx, client, piaTokenURL, username, password) + client.CloseIdleConnections() + cleanupTemporaryConnection(ctx, removeConnection, &err) + if err == nil { + return token, nil + } + errs = append(errs, err) + } + + return "", fmt.Errorf("fetching from token server: %w", errors.Join(errs...)) +} + +func fetchAddKey(ctx context.Context, client *http.Client, serverName string, + serverPort uint16, token, publicKey string, +) (data addKeyResponse, err error) { + const timeout = 10 * time.Second + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + query := make(url.Values) + query.Add("pt", token) + query.Add("pubkey", publicKey) + requestURL := url.URL{ + Scheme: "https", + Host: net.JoinHostPort(serverName, strconv.Itoa(int(serverPort))), + Path: "/addKey", + RawQuery: query.Encode(), + } + + request, err := http.NewRequestWithContext(ctx, http.MethodGet, requestURL.String(), nil) + if err != nil { + return data, replaceInErr(err, map[string]string{url.QueryEscape(token): ""}) + } + + response, err := client.Do(request) + if err != nil { + return data, replaceInErr(err, map[string]string{url.QueryEscape(token): ""}) + } + defer response.Body.Close() + + if response.StatusCode != http.StatusOK { + return data, makeNOKStatusError(response, + map[string]string{url.QueryEscape(token): ""}) + } + + err = json.NewDecoder(response.Body).Decode(&data) + if err != nil { + return data, fmt.Errorf("decoding response: %w", err) + } + if data.Status != "OK" { + return data, fmt.Errorf("bad response received with status %q", data.Status) + } + return data, nil +} + +func mapAddKeyResponse(connection models.Connection, privateKey string, + response addKeyResponse, +) (wireguardConnection models.WireguardConnection, err error) { + parsedPrivateKey, err := wgtypes.ParseKey(privateKey) + if err != nil { + return wireguardConnection, fmt.Errorf("client private key is not valid: %w", err) + } + registeredPublicKey := parsedPrivateKey.PublicKey().String() + if response.PeerPubKey != "" && response.PeerPubKey != registeredPublicKey { + return wireguardConnection, errors.New("registered client public key does not match generated private key") + } + + _, err = wgtypes.ParseKey(response.ServerKey) + if err != nil { + return wireguardConnection, fmt.Errorf("server public key is not valid: %w", err) + } + if response.ServerPort == 0 { + return wireguardConnection, errors.New("server port is not set") + } + err = validateRegistrationIPv4(response.ServerIP, "server IP") + if err != nil { + return wireguardConnection, err + } + err = validateRegistrationIPv4(response.ServerVIP, "server virtual IP") + if err != nil { + return wireguardConnection, err + } + err = validateRegistrationIPv4(response.PeerIP, "peer IP") + if err != nil { + return wireguardConnection, err + } + + connection.Type = vpn.Wireguard + connection.IP = response.ServerIP + connection.Port = response.ServerPort + connection.Protocol = constants.UDP + connection.PubKey = response.ServerKey + + address := netip.PrefixFrom(response.PeerIP, response.PeerIP.BitLen()) + return models.WireguardConnection{ + Connection: connection, + PrivateKey: privateKey, + Addresses: []netip.Prefix{address}, + DNSServers: append([]netip.Addr(nil), response.DNSServers...), + Gateway: response.ServerVIP, + }, nil +} + +func validateRegistrationIPv4(address netip.Addr, fieldName string) error { + switch { + case !address.IsValid(): + return fmt.Errorf("%s is not set", fieldName) + case address.IsUnspecified(): + return fmt.Errorf("%s is unspecified", fieldName) + case !address.Is4(): + return fmt.Errorf("%s is IPv6, which PIA registration does not support", fieldName) + default: + return nil + } +} diff --git a/internal/provider/privateinternetaccess/wireguard_test.go b/internal/provider/privateinternetaccess/wireguard_test.go new file mode 100644 index 000000000..f704062ad --- /dev/null +++ b/internal/provider/privateinternetaccess/wireguard_test.go @@ -0,0 +1,344 @@ +package privateinternetaccess + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "strconv" + "strings" + "testing" + "time" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/qdm12/gluetun/internal/models" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +func Test_Provider_RegisterWireguard(t *testing.T) { + t.Parallel() + + serverPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + serverPublicKey := serverPrivateKey.PublicKey().String() + client := &http.Client{ + Transport: piaRoundTripFunc(func(request *http.Request) (*http.Response, error) { + var body string + switch request.URL.Path { + case "/api/client/v2/token": + body = `{"token":"pia-token"}` + case "/addKey": + registeredPublicKey := request.URL.Query().Get("pubkey") + body = `{"status":"OK","server_key":"` + serverPublicKey + + `","server_port":51820,"server_ip":"198.51.100.3",` + + `"server_vip":"10.13.161.1","peer_ip":"10.13.161.2",` + + `"peer_pubkey":"` + registeredPublicKey + `",` + + `"dns_servers":["10.0.0.242"]}` + default: + t.Fatalf("unexpected request path %s", request.URL.Path) + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }), + } + provider := New(nil, time.Now, nil) + tokenServerIP := netip.MustParseAddr("192.0.2.20") + selectedServerIP := netip.MustParseAddr("198.51.100.2") + dialedIPs := make([]netip.Addr, 0, 2) + provider.newDialingClient = func(_ string, serverIP netip.Addr, + _ func(context.Context, string, string) (net.Conn, error), + ) (*http.Client, error) { + dialedIPs = append(dialedIPs, serverIP) + return client, nil + } + lookupNetIP := func(_ context.Context, network, host string) ([]netip.Addr, error) { + assert.Equal(t, "ip4", network) + assert.Equal(t, "www.privateinternetaccess.com", host) + return []netip.Addr{tokenServerIP}, nil + } + allowedConnections := make([]models.Connection, 0, 2) + removedAllowances := 0 + allowConnection := func(_ context.Context, connection models.Connection) ( + func(context.Context) error, error, + ) { + allowedConnections = append(allowedConnections, connection) + return func(context.Context) error { + removedAllowances++ + return nil + }, nil + } + selectedConnection := models.Connection{ + Type: vpn.Wireguard, + IP: selectedServerIP, + Port: wireguardRegistrationPort, + Protocol: constants.UDP, + Hostname: "ca-vancouver.privacy.network", + ServerName: testPIAServerName, + PortForward: true, + } + + registration, err := provider.RegisterWireguard(context.Background(), selectedConnection, + "username", "password", lookupNetIP, new(net.Dialer).DialContext, allowConnection) + require.NoError(t, err) + + expectedRegistrationConnection := selectedConnection + expectedRegistrationConnection.Protocol = constants.TCP + assert.Equal(t, []models.Connection{ + {IP: tokenServerIP, Port: 443, Protocol: constants.TCP}, + expectedRegistrationConnection, + }, allowedConnections) + assert.Equal(t, []netip.Addr{tokenServerIP, selectedServerIP}, dialedIPs) + assert.Equal(t, len(allowedConnections), removedAllowances) + assert.Equal(t, netip.MustParseAddr("198.51.100.3"), registration.Connection.IP) + assert.Equal(t, uint16(51820), registration.Connection.Port) + assert.Equal(t, constants.UDP, registration.Connection.Protocol) +} + +func Test_fetchAddKey(t *testing.T) { + t.Parallel() + + const serverName = testPIAServerName + const token = "pia-token" + clientPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + clientPublicKey := clientPrivateKey.PublicKey().String() + serverPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + serverPublicKey := serverPrivateKey.PublicKey().String() + + type receivedRequest struct { + host string + path string + token string + publicKey string + pk string + serverName string + } + received := make(chan receivedRequest, 1) + handler := http.HandlerFunc(func(responseWriter http.ResponseWriter, request *http.Request) { + received <- receivedRequest{ + host: request.Host, + path: request.URL.Path, + token: request.URL.Query().Get("pt"), + publicKey: request.URL.Query().Get("pubkey"), + pk: request.URL.Query().Get("pk"), + serverName: request.TLS.ServerName, + } + responseWriter.Header().Set("Content-Type", "application/json") + _, _ = responseWriter.Write([]byte(`{"status":"OK","server_key":"` + serverPublicKey + + `","server_port":1337,"server_ip":"1.2.3.4","server_vip":"1.2.3.5",` + + `"peer_ip":"10.13.161.2","dns_servers":["10.0.0.242"]}`)) + }) + server := httptest.NewUnstartedServer(handler) + + certificate, rootCAs := newTestServerCertificate(t, serverName) + server.TLS = &tls.Config{ + MinVersion: tls.VersionTLS12, + Certificates: []tls.Certificate{certificate}, + } + server.StartTLS() + t.Cleanup(server.Close) + + listenerAddress, err := netip.ParseAddrPort(server.Listener.Addr().String()) + require.NoError(t, err) + client, err := newHTTPClientDialing(serverName, listenerAddress.Addr(), new(net.Dialer).DialContext) + require.NoError(t, err) + transport, ok := client.Transport.(*http.Transport) + require.True(t, ok) + transport.TLSClientConfig.RootCAs = rootCAs + + data, err := fetchAddKey(context.Background(), client, serverName, + listenerAddress.Port(), token, clientPublicKey) + require.NoError(t, err) + + request := <-received + assert.Equal(t, net.JoinHostPort(serverName, strconv.Itoa(int(listenerAddress.Port()))), request.host) + assert.Equal(t, "/addKey", request.path) + assert.Equal(t, token, request.token) + assert.Equal(t, clientPublicKey, request.publicKey) + assert.Empty(t, request.pk) + assert.Equal(t, serverName, request.serverName) + assert.Equal(t, serverPublicKey, data.ServerKey) +} + +func Test_mapAddKeyResponse(t *testing.T) { + t.Parallel() + + serverPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + serverPublicKey := serverPrivateKey.PublicKey().String() + const registrationPort = 1337 + selectedConnection := models.Connection{ + Type: vpn.Wireguard, + IP: netip.MustParseAddr("198.51.100.2"), + Port: registrationPort, + Protocol: constants.UDP, + Hostname: "ca-vancouver.privacy.network", + ServerName: testPIAServerName, + PortForward: true, + } + const wireguardPort = 51820 + response := addKeyResponse{ + Status: "OK", + ServerKey: serverPublicKey, + ServerPort: wireguardPort, + ServerIP: netip.MustParseAddr("198.51.100.3"), + ServerVIP: netip.MustParseAddr("10.13.161.1"), + PeerIP: netip.MustParseAddr("10.13.161.2"), + DNSServers: []netip.Addr{netip.MustParseAddr("10.0.0.242")}, + } + + clientPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + response.PeerPubKey = clientPrivateKey.PublicKey().String() + registration, err := mapAddKeyResponse(selectedConnection, clientPrivateKey.String(), response) + require.NoError(t, err) + + expectedConnection := selectedConnection + expectedConnection.IP = response.ServerIP + expectedConnection.Port = response.ServerPort + expectedConnection.PubKey = response.ServerKey + assert.Equal(t, expectedConnection, registration.Connection) + assert.Equal(t, clientPrivateKey.String(), registration.PrivateKey) + assert.Equal(t, []netip.Prefix{netip.MustParsePrefix("10.13.161.2/32")}, registration.Addresses) + assert.Equal(t, response.DNSServers, registration.DNSServers) + assert.Equal(t, response.ServerVIP, registration.Gateway) + + response.PeerPubKey = serverPublicKey + _, err = mapAddKeyResponse(selectedConnection, clientPrivateKey.String(), response) + assert.ErrorContains(t, err, "registered client public key does not match generated private key") +} + +func Test_mapAddKeyResponse_rejectsUnsupportedIPAddresses(t *testing.T) { + t.Parallel() + + serverPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + clientPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + + validResponse := addKeyResponse{ + ServerKey: serverPrivateKey.PublicKey().String(), + ServerPort: 51820, + ServerIP: netip.MustParseAddr("198.51.100.3"), + ServerVIP: netip.MustParseAddr("10.13.161.1"), + PeerIP: netip.MustParseAddr("10.13.161.2"), + } + testCases := map[string]struct { + update func(response *addKeyResponse) + errMessage string + }{ + "server_ip_unspecified": { + update: func(response *addKeyResponse) { + response.ServerIP = netip.IPv4Unspecified() + }, + errMessage: "server IP is unspecified", + }, + "server_ip_ipv6": { + update: func(response *addKeyResponse) { + response.ServerIP = netip.MustParseAddr("2001:db8::3") + }, + errMessage: "server IP is IPv6, which PIA registration does not support", + }, + "server_virtual_ip_unspecified": { + update: func(response *addKeyResponse) { + response.ServerVIP = netip.IPv4Unspecified() + }, + errMessage: "server virtual IP is unspecified", + }, + "server_virtual_ip_ipv6": { + update: func(response *addKeyResponse) { + response.ServerVIP = netip.MustParseAddr("2001:db8::1") + }, + errMessage: "server virtual IP is IPv6, which PIA registration does not support", + }, + "peer_ip_unspecified": { + update: func(response *addKeyResponse) { + response.PeerIP = netip.IPv4Unspecified() + }, + errMessage: "peer IP is unspecified", + }, + "peer_ip_ipv6": { + update: func(response *addKeyResponse) { + response.PeerIP = netip.MustParseAddr("2001:db8::2") + }, + errMessage: "peer IP is IPv6, which PIA registration does not support", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + response := validResponse + testCase.update(&response) + _, err := mapAddKeyResponse(models.Connection{}, clientPrivateKey.String(), response) + + require.Error(t, err) + assert.ErrorContains(t, err, testCase.errMessage) + }) + } +} + +func newTestServerCertificate(t *testing.T, serverName string) ( + certificate tls.Certificate, rootCAs *x509.CertPool, +) { + t.Helper() + + caPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + now := time.Now() + const caSerialNumber = 1 + caTemplate := &x509.Certificate{ + SerialNumber: big.NewInt(caSerialNumber), + Subject: pkix.Name{CommonName: "PIA test CA"}, + NotBefore: now.Add(-time.Hour), + NotAfter: now.Add(time.Hour), + IsCA: true, + BasicConstraintsValid: true, + KeyUsage: x509.KeyUsageCertSign, + } + caDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, + &caPrivateKey.PublicKey, caPrivateKey) + require.NoError(t, err) + caCertificate, err := x509.ParseCertificate(caDER) + require.NoError(t, err) + + serverPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + const serverSerialNumber = 2 + serverTemplate := &x509.Certificate{ + SerialNumber: big.NewInt(serverSerialNumber), + Subject: pkix.Name{CommonName: serverName}, + DNSNames: []string{serverName}, + NotBefore: now.Add(-time.Hour), + NotAfter: now.Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + serverDER, err := x509.CreateCertificate(rand.Reader, serverTemplate, + caCertificate, &serverPrivateKey.PublicKey, caPrivateKey) + require.NoError(t, err) + + rootCAs = x509.NewCertPool() + rootCAs.AddCert(caCertificate) + return tls.Certificate{ + Certificate: [][]byte{serverDER, caDER}, + PrivateKey: serverPrivateKey, + }, rootCAs +} diff --git a/internal/provider/provider.go b/internal/provider/provider.go index d9b0bebb9..c85cf3930 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -2,6 +2,8 @@ package provider import ( "context" + "net" + "net/netip" "github.com/qdm12/gluetun/internal/configuration/settings" "github.com/qdm12/gluetun/internal/models" @@ -15,3 +17,20 @@ type Provider interface { FetchServers(ctx context.Context, minServers int) ( servers []models.Server, err error) } + +// DynamicWireguardProvider obtains a live server, registers a Wireguard key +// and returns provider-generated connection settings. Discovery and +// registration are done for every connection attempt. +type DynamicWireguardProvider interface { + GetWireguardConnection(ctx context.Context, selection settings.ServerSelection, + lookupNetIP func(context.Context, string, string) ([]netip.Addr, error), + dialContext func(context.Context, string, string) (net.Conn, error), + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error)) ( + connection models.Connection, err error) + RegisterWireguard(ctx context.Context, connection models.Connection, + username, password string, + lookupNetIP func(context.Context, string, string) ([]netip.Addr, error), + dialContext func(context.Context, string, string) (net.Conn, error), + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error)) ( + wireguardConnection models.WireguardConnection, err error) +} diff --git a/internal/storage/formatting.go b/internal/storage/formatting.go index 44865a5b1..9aceaeace 100644 --- a/internal/storage/formatting.go +++ b/internal/storage/formatting.go @@ -19,11 +19,13 @@ func noServerFoundError(selection settings.ServerSelection) (err error) { messageParts = append(messageParts, "VPN "+selection.VPN) - protocol := constants.UDP - if selection.OpenVPN.Protocol == constants.TCP { - protocol = constants.TCP + if selection.VPN == vpn.OpenVPN { + protocol := constants.UDP + if selection.OpenVPN.Protocol == constants.TCP { + protocol = constants.TCP + } + messageParts = append(messageParts, "protocol "+protocol) } - messageParts = append(messageParts, "protocol "+protocol) switch len(selection.Countries) { case 0: @@ -113,7 +115,7 @@ func noServerFoundError(selection settings.ServerSelection) (err error) { messageParts = append(messageParts, part) } - if *selection.OpenVPN.PIAEncPreset != "" { + if selection.VPN == vpn.OpenVPN && *selection.OpenVPN.PIAEncPreset != "" { part := "encryption preset " + *selection.OpenVPN.PIAEncPreset messageParts = append(messageParts, part) } @@ -150,7 +152,7 @@ func noServerFoundError(selection settings.ServerSelection) (err error) { if selection.VPN == vpn.Wireguard { targetIP = selection.Wireguard.EndpointIP } - if targetIP.IsValid() { + if targetIP.IsValid() && !targetIP.IsUnspecified() { messageParts = append(messageParts, "target ip address "+targetIP.String()) } diff --git a/internal/storage/formatting_test.go b/internal/storage/formatting_test.go new file mode 100644 index 000000000..1d30a3ab2 --- /dev/null +++ b/internal/storage/formatting_test.go @@ -0,0 +1,29 @@ +package storage + +import ( + "testing" + + "github.com/qdm12/gluetun/internal/configuration/settings" + "github.com/qdm12/gluetun/internal/constants/providers" + "github.com/qdm12/gluetun/internal/constants/vpn" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_noServerFoundError_wireguardOmitsOpenVPNSelection(t *testing.T) { + t.Parallel() + + selection := settings.ServerSelection{ + VPN: vpn.Wireguard, + Regions: []string{"CA Vancouver"}, + }.WithDefaults(providers.PrivateInternetAccess) + *selection.PortForwardOnly = true + + err := noServerFoundError(selection) + require.Error(t, err) + assert.ErrorContains(t, err, + "no server found: for VPN wireguard; region CA Vancouver; port forwarding only") + assert.NotContains(t, err.Error(), "protocol") + assert.NotContains(t, err.Error(), "encryption preset") + assert.NotContains(t, err.Error(), "target ip address") +} diff --git a/internal/vpn/bootstrap.go b/internal/vpn/bootstrap.go new file mode 100644 index 000000000..a589a246d --- /dev/null +++ b/internal/vpn/bootstrap.go @@ -0,0 +1,113 @@ +package vpn + +import ( + "context" + "errors" + "fmt" + "net" + "net/netip" + "strings" + "sync" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/models" +) + +type bootstrapResolver struct { + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error) + dialContext func(context.Context, string, string) (net.Conn, error) + lookupNetIP func(context.Context, string, string, + func(context.Context, string, string) (net.Conn, error)) ([]netip.Addr, error) +} + +func newBootstrapResolver( + mark uint32, + allowConnection func(context.Context, models.Connection) (func(context.Context) error, error), +) *bootstrapResolver { + return &bootstrapResolver{ + allowConnection: allowConnection, + dialContext: newPhysicalDialContext(mark), + lookupNetIP: func(ctx context.Context, network, host string, + dial func(context.Context, string, string) (net.Conn, error), + ) ([]netip.Addr, error) { + resolver := &net.Resolver{ + PreferGo: true, + Dial: dial, + } + return resolver.LookupNetIP(ctx, network, host) + }, + } +} + +func (r *bootstrapResolver) LookupNetIP(ctx context.Context, network, host string) ( + addresses []netip.Addr, err error, +) { + type allowedConnection struct { + remove func(context.Context) error + } + + allowedConnections := make([]allowedConnection, 0, 2) + allowedConnectionsSet := make(map[models.Connection]struct{}, 2) + var allowancesMutex sync.Mutex + dial := func(dialCtx context.Context, dialNetwork, address string) (net.Conn, error) { + connection, err := dnsConnection(dialNetwork, address) + if err != nil { + return nil, err + } + + allowancesMutex.Lock() + _, alreadyAllowed := allowedConnectionsSet[connection] + if !alreadyAllowed { + remove, err := r.allowConnection(dialCtx, connection) + if err != nil { + allowancesMutex.Unlock() + return nil, fmt.Errorf("allowing bootstrap DNS connection: %w", err) + } + allowedConnections = append(allowedConnections, allowedConnection{ + remove: remove, + }) + allowedConnectionsSet[connection] = struct{}{} + } + allowancesMutex.Unlock() + + return r.dialContext(dialCtx, dialNetwork, address) + } + + addresses, lookupErr := r.lookupNetIP(ctx, network, host, dial) + allowancesMutex.Lock() + allowedConnections = append([]allowedConnection(nil), allowedConnections...) + allowancesMutex.Unlock() + cleanupCtx := context.WithoutCancel(ctx) + cleanupErrs := make([]error, 0, len(allowedConnections)) + for i := len(allowedConnections) - 1; i >= 0; i-- { + removeErr := allowedConnections[i].remove(cleanupCtx) + if removeErr != nil { + cleanupErrs = append(cleanupErrs, + fmt.Errorf("removing bootstrap DNS connection: %w", removeErr)) + } + } + + return addresses, errors.Join(lookupErr, errors.Join(cleanupErrs...)) +} + +func dnsConnection(network, address string) (connection models.Connection, err error) { + switch { + case strings.HasPrefix(network, constants.UDP): + connection.Protocol = constants.UDP + case strings.HasPrefix(network, constants.TCP): + connection.Protocol = constants.TCP + default: + return connection, fmt.Errorf("DNS network is not supported: %s", network) + } + + addressPort, err := netip.ParseAddrPort(address) + if err != nil { + return connection, fmt.Errorf("parsing DNS resolver address: %w", err) + } + connection.IP = addressPort.Addr().Unmap() + if connection.IP.Is6() { + connection.IP = connection.IP.WithZone("") + } + connection.Port = addressPort.Port() + return connection, nil +} diff --git a/internal/vpn/bootstrap_dial_linux.go b/internal/vpn/bootstrap_dial_linux.go new file mode 100644 index 000000000..2e9783a2d --- /dev/null +++ b/internal/vpn/bootstrap_dial_linux.go @@ -0,0 +1,32 @@ +//go:build linux + +package vpn + +import ( + "context" + "errors" + "net" + "syscall" + "time" + + "golang.org/x/sys/unix" +) + +func newPhysicalDialContext(mark uint32) func(context.Context, string, string) (net.Conn, error) { + const connectionTimeout = 30 * time.Second + dialer := &net.Dialer{ + Timeout: connectionTimeout, + KeepAlive: connectionTimeout, + Control: func(_, _ string, rawConnection syscall.RawConn) error { + var setMarkErr error + controlErr := rawConnection.Control(func(fileDescriptor uintptr) { + // Wireguard's inverted policy rule ignores packets carrying its + // firewall mark, so these sockets use the physical main route. + setMarkErr = unix.SetsockoptInt(int(fileDescriptor), + unix.SOL_SOCKET, unix.SO_MARK, int(mark)) + }) + return errors.Join(controlErr, setMarkErr) + }, + } + return dialer.DialContext +} diff --git a/internal/vpn/bootstrap_dial_other.go b/internal/vpn/bootstrap_dial_other.go new file mode 100644 index 000000000..fbc6eb494 --- /dev/null +++ b/internal/vpn/bootstrap_dial_other.go @@ -0,0 +1,18 @@ +//go:build !linux + +package vpn + +import ( + "context" + "net" + "time" +) + +func newPhysicalDialContext(_ uint32) func(context.Context, string, string) (net.Conn, error) { + const connectionTimeout = 30 * time.Second + dialer := &net.Dialer{ + Timeout: connectionTimeout, + KeepAlive: connectionTimeout, + } + return dialer.DialContext +} diff --git a/internal/vpn/bootstrap_test.go b/internal/vpn/bootstrap_test.go new file mode 100644 index 000000000..ac50abb61 --- /dev/null +++ b/internal/vpn/bootstrap_test.go @@ -0,0 +1,122 @@ +package vpn + +import ( + "context" + "net" + "net/netip" + "testing" + + "github.com/qdm12/gluetun/internal/constants" + "github.com/qdm12/gluetun/internal/models" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func Test_bootstrapResolver_LookupNetIP(t *testing.T) { + t.Parallel() + + ctx := context.Background() + const resolverAddress = "100.100.100.100:53" + const host = "serverlist.piaservers.net" + expectedAddress := netip.MustParseAddr("192.0.2.10") + expectedConnection := models.Connection{ + IP: netip.MustParseAddr("100.100.100.100"), + Port: 53, + Protocol: constants.UDP, + } + + var allowedConnection models.Connection + allowanceRemoved := false + resolver := &bootstrapResolver{ + allowConnection: func(_ context.Context, connection models.Connection) ( + func(context.Context) error, error, + ) { + allowedConnection = connection + return func(context.Context) error { + allowanceRemoved = true + return nil + }, nil + }, + dialContext: func(_ context.Context, network, address string) (net.Conn, error) { + assert.Equal(t, constants.UDP, network) + assert.Equal(t, resolverAddress, address) + clientConnection, serverConnection := net.Pipe() + t.Cleanup(func() { + _ = serverConnection.Close() + }) + return clientConnection, nil + }, + lookupNetIP: func(lookupCtx context.Context, network, lookupHost string, + dial func(context.Context, string, string) (net.Conn, error), + ) ([]netip.Addr, error) { + assert.Equal(t, ctx, lookupCtx) + assert.Equal(t, "ip4", network) + assert.Equal(t, host, lookupHost) + connection, err := dial(lookupCtx, constants.UDP, resolverAddress) + require.NoError(t, err) + require.NoError(t, connection.Close()) + assert.False(t, allowanceRemoved) + return []netip.Addr{expectedAddress}, nil + }, + } + + addresses, err := resolver.LookupNetIP(ctx, "ip4", host) + require.NoError(t, err) + assert.Equal(t, []netip.Addr{expectedAddress}, addresses) + assert.Equal(t, expectedConnection, allowedConnection) + assert.True(t, allowanceRemoved) +} + +func Test_dnsConnection(t *testing.T) { + t.Parallel() + + testCases := map[string]struct { + network string + address string + connection models.Connection + errMessage string + }{ + "udp_ipv4": { + network: constants.UDP + "4", + address: "100.100.100.100:53", + connection: models.Connection{ + IP: netip.MustParseAddr("100.100.100.100"), + Port: 53, + Protocol: constants.UDP, + }, + }, + "tcp_ipv6": { + network: constants.TCP + "6", + address: "[2001:db8::53]:53", + connection: models.Connection{ + IP: netip.MustParseAddr("2001:db8::53"), + Port: 53, + Protocol: constants.TCP, + }, + }, + "unsupported_network": { + network: "ip", + address: "100.100.100.100:53", + errMessage: "DNS network is not supported", + }, + "invalid_address": { + network: constants.UDP, + address: "invalid", + errMessage: "parsing DNS resolver address", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + connection, err := dnsConnection(testCase.network, testCase.address) + if testCase.errMessage != "" { + assert.ErrorContains(t, err, testCase.errMessage) + return + } + require.NoError(t, err) + assert.Equal(t, testCase.connection, connection) + }) + } +} diff --git a/internal/vpn/interfaces.go b/internal/vpn/interfaces.go index 74c128e50..0cbc48f67 100644 --- a/internal/vpn/interfaces.go +++ b/internal/vpn/interfaces.go @@ -17,6 +17,8 @@ import ( type Firewall interface { SetVPNConnection(ctx context.Context, connection models.Connection, interfaceName string) error + TempAllowConnection(ctx context.Context, connection models.Connection) ( + remove func(context.Context) error, err error) SetAllowedPort(ctx context.Context, port uint16, interfaceName string) error RemoveAllowedPort(ctx context.Context, port uint16) error tcp.Firewall diff --git a/internal/vpn/portforward.go b/internal/vpn/portforward.go index ec8a7862c..ec5b24817 100644 --- a/internal/vpn/portforward.go +++ b/internal/vpn/portforward.go @@ -28,6 +28,7 @@ func (l *Loop) startPortForwarding(data tunnelUpData) (err error) { Service: service.Settings{ PortForwarder: data.portForwarder, Interface: data.vpnIntf, + Gateway: data.gateway, ServerName: data.serverName, CanPortForward: data.canPortForward, Username: data.username, diff --git a/internal/vpn/run.go b/internal/vpn/run.go index 1e49bbe43..b59f4366a 100644 --- a/internal/vpn/run.go +++ b/internal/vpn/run.go @@ -2,6 +2,7 @@ package vpn import ( "context" + "net/netip" "github.com/qdm12/gluetun/internal/constants" "github.com/qdm12/gluetun/internal/constants/vpn" @@ -31,6 +32,7 @@ func (l *Loop) Run(ctx context.Context, done chan<- struct{}) { } var vpnInterface string var connection models.Connection + var gateway netip.Addr var err error subLogger := l.logger.New(log.SetComponent(settings.Type)) switch settings.Type { @@ -44,7 +46,7 @@ func (l *Loop) Run(ctx context.Context, done chan<- struct{}) { l.openvpnConf, providerConf, settings, l.ipv6SupportLevel, l.cmder, subLogger) case vpn.Wireguard: vpnInterface = settings.Wireguard.Interface - vpnRunner, connection, err = setupWireguard(ctx, l.netLinker, l.fw, + vpnRunner, connection, gateway, err = setupWireguard(ctx, l.netLinker, l.fw, providerConf, settings, l.ipv6SupportLevel, subLogger) default: panic("vpn type not implemented: " + settings.Type) @@ -53,6 +55,10 @@ func (l *Loop) Run(ctx context.Context, done chan<- struct{}) { l.crashed(ctx, err) continue } + serverName := connection.ServerName + if settings.Provider.PortForwarding.ServerName != "" { + serverName = settings.Provider.PortForwarding.ServerName + } tunnelUpData := tunnelUpData{ upCommand: *settings.UpCommand, pmtud: tunnelUpPMTUDData{ @@ -64,10 +70,11 @@ func (l *Loop) Run(ctx context.Context, done chan<- struct{}) { tcpAddrs: settings.PMTUD.TCPAddresses, }, serverIP: connection.IP, - serverName: connection.ServerName, + serverName: serverName, canPortForward: connection.PortForward, portForwarder: portForwarder, vpnIntf: vpnInterface, + gateway: gateway, username: settings.Provider.PortForwarding.Username, password: settings.Provider.PortForwarding.Password, } diff --git a/internal/vpn/tunnelup.go b/internal/vpn/tunnelup.go index 4843812ab..4ab2b5cf8 100644 --- a/internal/vpn/tunnelup.go +++ b/internal/vpn/tunnelup.go @@ -25,6 +25,7 @@ type tunnelUpData struct { pmtud tunnelUpPMTUDData // Port forwarding vpnIntf string + gateway netip.Addr serverName string // used for PIA canPortForward bool // used for PIA username string // used for PIA diff --git a/internal/vpn/wireguard.go b/internal/vpn/wireguard.go index 2fe4b15c4..9f2e26594 100644 --- a/internal/vpn/wireguard.go +++ b/internal/vpn/wireguard.go @@ -3,7 +3,10 @@ package vpn import ( "context" "fmt" + "net" "net/netip" + "strings" + "time" "github.com/qdm12/gluetun/internal/configuration/settings" "github.com/qdm12/gluetun/internal/models" @@ -11,22 +14,67 @@ import ( "github.com/qdm12/gluetun/internal/provider" "github.com/qdm12/gluetun/internal/wireguard" "github.com/qdm12/gosettings" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" ) // setupWireguard sets Wireguard up using the configurators and settings given. -// It returns a serverName for port forwarding (PIA) and an error if it fails. +// It returns the selected connection, an optional provider-generated gateway +// for port forwarding, and an error if setup fails. func setupWireguard(ctx context.Context, netlinker NetLinker, fw Firewall, providerConf provider.Provider, settings settings.VPN, ipv6SupportLevel netlink.IPv6SupportLevel, logger wireguard.Logger) ( - wireguarder *wireguard.Wireguard, connection models.Connection, err error, + wireguarder *wireguard.Wireguard, connection models.Connection, + gateway netip.Addr, err error, ) { ipv6Internet := ipv6SupportLevel == netlink.IPv6Internet - connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet) + dynamicProvider, dynamic := providerConf.(provider.DynamicWireguardProvider) + var lookupNetIP func(context.Context, string, string) ([]netip.Addr, error) + var bootstrapDialContext func(context.Context, string, string) (net.Conn, error) + if dynamic { + bootstrapSettings := wireguard.Settings{} + bootstrapSettings.SetDefaults() + bootstrapResolver := newBootstrapResolver(bootstrapSettings.FirewallMark, + fw.TempAllowConnection) + lookupNetIP = bootstrapResolver.LookupNetIP + bootstrapDialContext = bootstrapResolver.dialContext + connection, err = dynamicProvider.GetWireguardConnection(ctx, + settings.Provider.ServerSelection, lookupNetIP, bootstrapDialContext, + fw.TempAllowConnection) + } else { + connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet) + } if err != nil { - return nil, models.Connection{}, fmt.Errorf("finding a VPN server: %w", err) + return nil, models.Connection{}, netip.Addr{}, fmt.Errorf("finding a VPN server: %w", err) } - wireguardSettings := buildWireguardSettings(connection, settings.Wireguard, ipv6SupportLevel.IsSupported()) + var wireguardSettings wireguard.Settings + if dynamic { + wireguardConnection, err := dynamicProvider.RegisterWireguard(ctx, connection, + *settings.OpenVPN.User, *settings.OpenVPN.Password, + lookupNetIP, bootstrapDialContext, fw.TempAllowConnection) + if err != nil { + return nil, models.Connection{}, netip.Addr{}, fmt.Errorf("registering Wireguard connection: %w", err) + } + connection = wireguardConnection.Connection + wireguardSettings = buildRegisteredWireguardSettings(wireguardConnection, + settings.Wireguard, ipv6SupportLevel.IsSupported()) + gateway = wireguardConnection.Gateway + clientPrivateKey, err := wgtypes.ParseKey(wireguardSettings.PrivateKey) + if err != nil { + return nil, models.Connection{}, netip.Addr{}, fmt.Errorf("parsing registered private key: %w", err) + } + logger.Info(fmt.Sprintf("PIA addKey Wireguard configuration: endpoint %s, peer public key %s, "+ + "client public key %s, interface addresses %s", + wireguardSettings.Endpoint, shortWireguardKey(wireguardSettings.PublicKey), + shortWireguardKey(clientPrivateKey.PublicKey().String()), + wireguardAddresses(wireguardSettings.Addresses))) + if len(wireguardConnection.DNSServers) > 0 { + logger.Debugf("Wireguard provider DNS servers: %v", wireguardConnection.DNSServers) + } + } else { + wireguardSettings = buildWireguardSettings(connection, + settings.Wireguard, ipv6SupportLevel.IsSupported()) + } logger.Debug("Wireguard server public key: " + wireguardSettings.PublicKey) logger.Debug("Wireguard client private key: " + gosettings.ObfuscateKey(wireguardSettings.PrivateKey)) @@ -34,15 +82,44 @@ func setupWireguard(ctx context.Context, netlinker NetLinker, wireguarder, err = wireguard.New(wireguardSettings, netlinker, logger) if err != nil { - return nil, models.Connection{}, fmt.Errorf("creating Wireguard: %w", err) + return nil, models.Connection{}, netip.Addr{}, fmt.Errorf("creating Wireguard: %w", err) } err = fw.SetVPNConnection(ctx, connection, settings.Wireguard.Interface) if err != nil { - return nil, models.Connection{}, fmt.Errorf("setting firewall: %w", err) + return nil, models.Connection{}, netip.Addr{}, fmt.Errorf("setting firewall: %w", err) } - return wireguarder, connection, nil + return wireguarder, connection, gateway, nil +} + +func buildRegisteredWireguardSettings(registration models.WireguardConnection, + userSettings settings.Wireguard, ipv6Supported bool, +) wireguard.Settings { + privateKey := registration.PrivateKey + userSettings.PrivateKey = &privateKey + userSettings.Addresses = append([]netip.Prefix(nil), registration.Addresses...) + if *userSettings.PersistentKeepaliveInterval == 0 { + piaPersistentKeepalive := 25 * time.Second + userSettings.PersistentKeepaliveInterval = &piaPersistentKeepalive + } + return buildWireguardSettings(registration.Connection, userSettings, ipv6Supported) +} + +func shortWireguardKey(key string) string { + const visibleCharacters = 8 + if len(key) <= 2*visibleCharacters { + return key + } + return key[:visibleCharacters] + "..." + key[len(key)-visibleCharacters:] +} + +func wireguardAddresses(addresses []netip.Prefix) string { + addressStrings := make([]string, len(addresses)) + for i, address := range addresses { + addressStrings[i] = address.String() + } + return strings.Join(addressStrings, ", ") } func buildWireguardSettings(connection models.Connection, diff --git a/internal/vpn/wireguard_test.go b/internal/vpn/wireguard_test.go index 074179232..a1bd96939 100644 --- a/internal/vpn/wireguard_test.go +++ b/internal/vpn/wireguard_test.go @@ -9,8 +9,46 @@ import ( "github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/wireguard" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" ) +func Test_buildRegisteredWireguardSettings(t *testing.T) { + t.Parallel() + + registeredPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + serverPrivateKey, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + registration := models.WireguardConnection{ + Connection: models.Connection{ + IP: netip.MustParseAddr("198.51.100.3"), + Port: 1337, + PubKey: serverPrivateKey.PublicKey().String(), + }, + PrivateKey: registeredPrivateKey.String(), + Addresses: []netip.Prefix{netip.MustParsePrefix("10.13.161.2/32")}, + } + zeroKeepalive := time.Duration(0) + userSettings := settings.Wireguard{ + PrivateKey: new("unused-user-private-key"), + PreSharedKey: new(""), + AllowedIPs: []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0")}, + PersistentKeepaliveInterval: &zeroKeepalive, + Interface: "tun0", + MTU: new(uint32(1320)), + } + + wireguardSettings := buildRegisteredWireguardSettings(registration, userSettings, false) + + assert.Equal(t, registeredPrivateKey.String(), wireguardSettings.PrivateKey) + assert.Equal(t, registration.Connection.PubKey, wireguardSettings.PublicKey) + assert.Equal(t, netip.MustParseAddrPort("198.51.100.3:1337"), wireguardSettings.Endpoint) + assert.Equal(t, registration.Addresses, wireguardSettings.Addresses) + assert.Equal(t, []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0")}, wireguardSettings.AllowedIPs) + assert.Equal(t, 25*time.Second, wireguardSettings.PersistentKeepaliveInterval) +} + func Test_buildWireguardSettings(t *testing.T) { t.Parallel() diff --git a/internal/wireguard/config.go b/internal/wireguard/config.go index 9dd3de74c..cddae9dd6 100644 --- a/internal/wireguard/config.go +++ b/internal/wireguard/config.go @@ -52,6 +52,13 @@ func makeDeviceConfig(settings Settings) (config wgtypes.Config, err error) { } firewallMark := int(settings.FirewallMark) + allowedIPs := make([]net.IPNet, len(settings.AllowedIPs)) + for i, allowedIP := range settings.AllowedIPs { + allowedIPs[i] = net.IPNet{ + IP: allowedIP.Addr().AsSlice(), + Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), + } + } config = wgtypes.Config{ PrivateKey: &privateKey, @@ -59,18 +66,9 @@ func makeDeviceConfig(settings Settings) (config wgtypes.Config, err error) { FirewallMark: &firewallMark, Peers: []wgtypes.PeerConfig{ { - PublicKey: publicKey, - PresharedKey: preSharedKey, - AllowedIPs: []net.IPNet{ - { - IP: net.IPv4(0, 0, 0, 0), - Mask: []byte{0, 0, 0, 0}, - }, - { - IP: net.IPv6zero, - Mask: []byte(net.IPv6zero), - }, - }, + PublicKey: publicKey, + PresharedKey: preSharedKey, + AllowedIPs: allowedIPs, PersistentKeepaliveInterval: persistentKeepaliveInterval, ReplaceAllowedIPs: true, Endpoint: &net.UDPAddr{ diff --git a/internal/wireguard/config_test.go b/internal/wireguard/config_test.go index cff564d26..82caa7b68 100644 --- a/internal/wireguard/config_test.go +++ b/internal/wireguard/config_test.go @@ -62,6 +62,7 @@ func Test_makeDeviceConfig(t *testing.T) { PreSharedKey: validKey3, FirewallMark: 9876, Endpoint: netip.AddrPortFrom(netip.AddrFrom4([4]byte{99, 99, 99, 99}), 51820), + AllowedIPs: []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0")}, }, config: wgtypes.Config{ PrivateKey: parseKey(t, validKey1), @@ -73,13 +74,9 @@ func Test_makeDeviceConfig(t *testing.T) { PresharedKey: parseKey(t, validKey3), AllowedIPs: []net.IPNet{ { - IP: net.IPv4(0, 0, 0, 0), + IP: net.IP{0, 0, 0, 0}, Mask: []byte{0, 0, 0, 0}, }, - { - IP: net.IPv6zero, - Mask: []byte(net.IPv6zero), - }, }, ReplaceAllowedIPs: true, Endpoint: &net.UDPAddr{ From c2f70b8efea8135f188ad809a3d7d185da15603a Mon Sep 17 00:00:00 2001 From: Neil Date: Mon, 13 Jul 2026 09:55:30 -0600 Subject: [PATCH 3/3] feat(healthcheck): retry the startup check within a budget The VPN startup healthcheck performed a single 6s TCP+TLS check. When a tunnel's DNS server takes a few seconds to become ready after connect (e.g. providers that do pre-tunnel work), that single attempt could fail and trigger an endless VPN restart loop. Retry the short (6s) parallel TCP+TLS attempts within a 60s total budget with a 2s backoff, returning on the first success and preserving the existing aggregated error on budget exhaustion. Periodic checks are unchanged. This removes the need for any manual startup-grace tuning and is a general robustness improvement for all providers. --- internal/healthcheck/checker.go | 94 +++++++++++++++---- internal/healthcheck/checker_test.go | 133 +++++++++++++++++++++++++++ 2 files changed, 208 insertions(+), 19 deletions(-) diff --git a/internal/healthcheck/checker.go b/internal/healthcheck/checker.go index 5d36f58d0..38295f1cf 100644 --- a/internal/healthcheck/checker.go +++ b/internal/healthcheck/checker.go @@ -67,16 +67,18 @@ func (c *Checker) SetConfig(tlsDialAddrs []string, icmpTargets []netip.Addr, // internal field startupOnFail, which is set by calling [Checker.SetConfig]. // // By default, startupOnFail should be false and the behavior is as follows: -// A blocking 6s-timed TCP+TLS check is performed first. If it fails, -// an error is returned and the [Checker] is not started. +// Blocking TCP+TLS checks are retried for up to 60 seconds, with each attempt +// limited to 6 seconds. If all attempts fail, an error is returned and the +// [Checker] is not started. // On success, it starts the periodic checks in a separate goroutine, returning // the runError error channel and a nil error. // // If startupOnFail is true, the behavior is as follows: -// A blocking 6s-timed TCP+TLS check is performed first. If it fails, -// the error is sent to the runError channel, but no error is returned -// and the [Checker] continues to start the periodic checks in a separate goroutine, returning -// the runError error channel and a nil error. +// Blocking TCP+TLS checks are retried for up to 60 seconds, with each attempt +// limited to 6 seconds. If all attempts fail, the error is sent to the runError +// channel, but no error is returned, and the [Checker] continues to start the +// periodic checks in a separate goroutine, returning the runError error channel +// and a nil error. // // The periodic checks consist in: // - a "small" ICMP echo check every minute @@ -213,6 +215,13 @@ func (c *Checker) fullPeriodicCheck(ctx context.Context) error { } func tcpTLSCheck(ctx context.Context, dialer *net.Dialer, targetAddress string) error { + return tcpTLSCheckWithDialContext(ctx, dialer.DialContext, targetAddress) +} + +func tcpTLSCheckWithDialContext(ctx context.Context, + dialContext func(context.Context, string, string) (net.Conn, error), + targetAddress string, +) (err error) { // TODO use mullvad API if current provider is Mullvad address, err := makeAddressToDial(targetAddress) @@ -221,10 +230,16 @@ func tcpTLSCheck(ctx context.Context, dialer *net.Dialer, targetAddress string) } const dialNetwork = "tcp4" - connection, err := dialer.DialContext(ctx, dialNetwork, address) + connection, err := dialContext(ctx, dialNetwork, address) if err != nil { return fmt.Errorf("dialing: %w", err) } + defer func() { + closeErr := connection.Close() + if err == nil && closeErr != nil { + err = fmt.Errorf("closing connection: %w", closeErr) + } + }() if strings.HasSuffix(address, ":443") { host, _, err := net.SplitHostPort(address) @@ -242,11 +257,6 @@ func tcpTLSCheck(ctx context.Context, dialer *net.Dialer, targetAddress string) } } - err = connection.Close() - if err != nil { - return fmt.Errorf("closing connection: %w", err) - } - return nil } @@ -298,17 +308,63 @@ func withRetries(ctx context.Context, tryTimeouts []time.Duration, return fmt.Errorf("all check tries failed:\n\t%s", strings.Join(errStrings, "\n\t")) } +// startupCheck retries short parallel TCP+TLS attempts within a longer startup +// budget so the tunnel and its DNS server have time to become ready. func (c *Checker) startupCheck(ctx context.Context) error { - // connection isn't under load yet when the checker starts, so a short - // 6 seconds timeout suffices and provides quick enough feedback that - // the new connection is not working. However, since the addresses to dial - // may be multiple, we run the check in parallel. If any succeeds, the check passes. + // Allow enough time for the VPN tunnel and its DNS server to become ready + // while retaining short individual checks for fast success and feedback. + const totalBudget = 60 * time.Second + const attemptTimeout = 6 * time.Second + const retryInterval = 2 * time.Second + return startupCheckWithRetries(ctx, totalBudget, attemptTimeout, retryInterval, c.startupCheckAttempt) +} + +func startupCheckWithRetries(ctx context.Context, totalBudget, attemptTimeout, retryInterval time.Duration, + check func(context.Context) error, +) error { + startupCtx, cancelStartup := context.WithTimeout(ctx, totalBudget) + defer cancelStartup() + + var lastErr error + for { + attemptCtx, cancelAttempt := context.WithTimeout(startupCtx, attemptTimeout) + lastErr = check(attemptCtx) + cancelAttempt() + if lastErr == nil { + return nil + } + + err := ctx.Err() + if err != nil { + return fmt.Errorf("checking startup connection: %w", err) + } + if startupCtx.Err() != nil { + return lastErr + } + + retryTimer := time.NewTimer(retryInterval) + select { + case <-ctx.Done(): + retryTimer.Stop() + return fmt.Errorf("checking startup connection: %w", ctx.Err()) + case <-startupCtx.Done(): + retryTimer.Stop() + return lastErr + case <-retryTimer.C: + } + } +} + +func (c *Checker) startupCheckAttempt(ctx context.Context) error { + // The connection isn't under load yet when the checker starts, so each short + // 6-second attempt provides quick feedback for the retry loop. Since the + // addresses to dial may be multiple, we run the check in parallel. If any + // succeeds, the check passes. // This is to prevent false negatives at startup, if one of the addresses is down // for external reasons. - const timeout = 6 * time.Second - ctx, cancel := context.WithTimeout(ctx, timeout) + ctx, cancel := context.WithCancel(ctx) defer cancel() - errCh := make(chan error) + errCh := make(chan error, len(c.tlsDialAddrs)) for _, address := range c.tlsDialAddrs { go func(addr string) { diff --git a/internal/healthcheck/checker_test.go b/internal/healthcheck/checker_test.go index f6241f768..bfa4f9820 100644 --- a/internal/healthcheck/checker_test.go +++ b/internal/healthcheck/checker_test.go @@ -2,6 +2,7 @@ package healthcheck import ( "context" + "errors" "net" "testing" "time" @@ -10,6 +11,84 @@ import ( "github.com/stretchr/testify/require" ) +func Test_startupCheckWithRetries(t *testing.T) { + t.Parallel() + + t.Run("success_on_later_retry", func(t *testing.T) { + t.Parallel() + + attempts := 0 + check := func(context.Context) error { + attempts++ + if attempts < 3 { + return errors.New("not ready") + } + return nil + } + + err := startupCheckWithRetries(t.Context(), time.Second, 100*time.Millisecond, time.Millisecond, check) + + assert.NoError(t, err) + assert.Equal(t, 3, attempts) + }) + + t.Run("total_budget_exhaustion", func(t *testing.T) { + t.Parallel() + + attempts := 0 + check := func(context.Context) error { + attempts++ + return errors.New("not ready") + } + const totalBudget = 30 * time.Millisecond + const minimumDuration = 25 * time.Millisecond + start := time.Now() + + err := startupCheckWithRetries(t.Context(), totalBudget, time.Second, 5*time.Millisecond, check) + + require.Error(t, err) + assert.ErrorContains(t, err, "not ready") + assert.Greater(t, attempts, 1) + assert.GreaterOrEqual(t, time.Since(start), minimumDuration) + }) + + t.Run("immediate_success_returns_fast", func(t *testing.T) { + t.Parallel() + + const maximumDuration = 250 * time.Millisecond + start := time.Now() + err := startupCheckWithRetries(t.Context(), time.Second, time.Second, 500*time.Millisecond, + func(context.Context) error { return nil }) + + assert.NoError(t, err) + assert.Less(t, time.Since(start), maximumDuration) + }) + + t.Run("context_cancellation_aborts_promptly", func(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(t.Context()) + checkStarted := make(chan struct{}) + check := func(ctx context.Context) error { + close(checkStarted) + <-ctx.Done() + return ctx.Err() + } + go func() { + <-checkStarted + cancel() + }() + const maximumDuration = 250 * time.Millisecond + start := time.Now() + + err := startupCheckWithRetries(ctx, time.Second, time.Second, 500*time.Millisecond, check) + + require.Error(t, err) + assert.ErrorIs(t, err, context.Canceled) + assert.Less(t, time.Since(start), maximumDuration) + }) +} + func Test_Checker_fullcheck(t *testing.T) { t.Parallel() @@ -98,3 +177,57 @@ func Test_makeAddressToDial(t *testing.T) { }) } } + +func Test_tcpTLSCheckWithDialContext_closesConnectionOnHandshakeError(t *testing.T) { + t.Parallel() + + connection := &failingTLSConnection{} + dialContext := func(_ context.Context, network, address string) (net.Conn, error) { + assert.Equal(t, "tcp4", network) + assert.Equal(t, "example.com:443", address) + return connection, nil + } + + err := tcpTLSCheckWithDialContext(t.Context(), dialContext, "example.com:443") + + require.Error(t, err) + assert.ErrorContains(t, err, "running TLS handshake") + assert.True(t, connection.closed) +} + +type failingTLSConnection struct { + closed bool +} + +func (c *failingTLSConnection) Read([]byte) (int, error) { + return 0, errors.New("TLS read error") +} + +func (c *failingTLSConnection) Write(buffer []byte) (int, error) { + return len(buffer), nil +} + +func (c *failingTLSConnection) Close() error { + c.closed = true + return nil +} + +func (c *failingTLSConnection) LocalAddr() net.Addr { + return &net.TCPAddr{} +} + +func (c *failingTLSConnection) RemoteAddr() net.Addr { + return &net.TCPAddr{} +} + +func (c *failingTLSConnection) SetDeadline(time.Time) error { + return nil +} + +func (c *failingTLSConnection) SetReadDeadline(time.Time) error { + return nil +} + +func (c *failingTLSConnection) SetWriteDeadline(time.Time) error { + return nil +}