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/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) + }) + } +} 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 +} 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{