Compare commits

..

17 Commits

Author SHA1 Message Date
Quentin McGaw 00d944e713 fix(protonvpn/updater): fallback to email if username is empty in auth info response 2026-05-21 16:58:33 +00:00
Quentin McGaw beda1764b1 feat(protonvpn): updater finds more servers using app-version linux-vpn 2026-05-21 16:51:36 +00:00
Quentin McGaw 2210a0e9ad fix(command): fix rare race condition on log line stream at command completion 2026-05-21 15:44:21 +00:00
Quentin McGaw f8a677a424 hotfix(portforward): log both external and internal ports when they diverge
- useful for ProtonVPN only
- clarify things up for the user
2026-05-21 14:45:40 +00:00
Quentin McGaw 8f012014d6 hotfix(firewall/iptables): only save stdout from iptables-save, not stderr 2026-05-21 03:50:44 +00:00
Quentin McGaw b119325241 hotfix(storage): do not write filepath field for non-manifest files 2026-05-19 03:03:30 +00:00
Quentin McGaw 7720b1fad4 fix(storage): ignore empty manifest servers file
- Fix #3318
2026-05-19 02:53:45 +00:00
Quentin McGaw 854bf5811d fix(wireguard): skip tun device checks when using kernelspace 2026-05-19 02:46:40 +00:00
Quentin McGaw 8f82376996 feat(storage): storage file structure changes (#3301)
- migrate persisted server data storage from `/gluetun/servers.json` to `/gluetun/servers/`
- add `STORAGE_SERVERS_ENABLED=on` to enable or disable on-disk server data storage
- add `STORAGE_SERVERS_DIRECTORY_PATH=/gluetun/servers` to configure where per-provider server files are stored
- keep backward compatibility with legacy `STORAGE_FILEPATH=/gluetun/servers.json`
- automatically read and migrate legacy `/gluetun/servers.json` into the new `/gluetun/servers/` layout when needed
- try to remove the legacy servers file after a successful migration to the new storage directory
- switch persisted server data from one large JSON file to a manifest plus per-provider JSON files
- add `UPDATER_PREFER_DIRECT_DOWNLOAD` to allow preferring direct download of provider server data
- keep deprecated updater flags `-enduser` and `-maintainer` as no-op warnings for backward compatibility
- preserve compatibility checks so persisted server data is discarded when its schema version no longer matches the built-in data
- allow preferred persisted provider data to override built-in data when versions match
- servers data now lives at https://github.com/qdm12/gluetun-servers/tree/main/pkg/servers
2026-05-19 04:28:25 +02:00
Immanuel Tikhonov cd19093d1d fix(openvpn/extract): trim spaces in config lines before parsing (#3327) 2026-05-12 03:44:29 +02:00
Quentin McGaw fd12e5f9e7 chore(provider/utils): fix flaky test caused by new random shuffle 2026-05-12 01:28:11 +00:00
Quentin McGaw 3ca4b48887 hotfix(provider/utils): randomize pool of filterd servers to pick connections from 2026-05-12 01:08:19 +00:00
Quentin McGaw 38cf094573 chore(boringpoll): remove gluetun.com which is now DOWN 🎉 2026-05-12 00:58:23 +00:00
Quentin McGaw 5b01324d5f hotfix(pmtud): detect IPv6 usage in VPN connection 2026-05-09 14:40:04 +00:00
Quentin McGaw 445f99d9dc hotfix(openvpn): bump hand-window from 10s to 20s 2026-05-08 16:12:13 +00:00
Quentin McGaw 891249849a fix(provider/pia): handle "port is busy" messages and retry port forwarding logic 2026-05-08 04:16:15 +00:00
Quentin McGaw 5cae870745 feat(provider/pia): try parsing JSON on bad port forwarding API status codes 2026-05-08 04:15:30 +00:00
68 changed files with 1249 additions and 304300 deletions
+3 -2
View File
@@ -249,6 +249,7 @@ ENV VPN_SERVICE_PROVIDER=pia \
UPDATER_PERIOD=0 \ UPDATER_PERIOD=0 \
UPDATER_MIN_RATIO=0.8 \ UPDATER_MIN_RATIO=0.8 \
UPDATER_VPN_SERVICE_PROVIDERS= \ UPDATER_VPN_SERVICE_PROVIDERS= \
UPDATER_PREFER_DIRECT_DOWNLOAD=no \
UPDATER_PROTONVPN_EMAIL= \ UPDATER_PROTONVPN_EMAIL= \
UPDATER_PROTONVPN_PASSWORD= \ UPDATER_PROTONVPN_PASSWORD= \
# Public IP # Public IP
@@ -257,7 +258,8 @@ ENV VPN_SERVICE_PROVIDER=pia \
PUBLICIP_API=ipinfo,ifconfigco,ip2location,cloudflare \ PUBLICIP_API=ipinfo,ifconfigco,ip2location,cloudflare \
PUBLICIP_API_TOKEN= \ PUBLICIP_API_TOKEN= \
# Storage # Storage
STORAGE_FILEPATH=/gluetun/servers.json \ STORAGE_SERVERS_ENABLED=on \
STORAGE_SERVERS_DIRECTORY_PATH=/gluetun/servers/ \
# Pprof # Pprof
PPROF_ENABLED=no \ PPROF_ENABLED=no \
PPROF_BLOCK_PROFILE_RATE=0 \ PPROF_BLOCK_PROFILE_RATE=0 \
@@ -265,7 +267,6 @@ ENV VPN_SERVICE_PROVIDER=pia \
PPROF_HTTP_SERVER_ADDRESS=":6060" \ PPROF_HTTP_SERVER_ADDRESS=":6060" \
# Extras # Extras
VERSION_INFORMATION=on \ VERSION_INFORMATION=on \
BORINGPOLL_GLUETUNCOM=off \
TZ= \ TZ= \
PUID=1000 \ PUID=1000 \
PGID=1000 PGID=1000
-4
View File
@@ -132,10 +132,6 @@ services:
[![Star History Chart](https://api.star-history.com/svg?repos=qdm12/gluetun&type=date&legend=top-left)](https://www.star-history.com/#qdm12/gluetun&type=date&legend=top-left) [![Star History Chart](https://api.star-history.com/svg?repos=qdm12/gluetun&type=date&legend=top-left)](https://www.star-history.com/#qdm12/gluetun&type=date&legend=top-left)
## Fight AI scamming
💁 You can optionally set `BORINGPOLL_GLUETUNCOM=on` to... [poll](./internal/boringpoll/boringpoll.go) that **scammy AI slop** website every few minutes so it costs them in terms of egress traffic. My gentle email reminders to take it down are being grossly ignored 🤷 This would make me very happy and serve this community.
## License ## License
[![MIT](https://img.shields.io/github/license/qdm12/gluetun)](https://github.com/qdm12/gluetun/blob/master/LICENSE) [![MIT](https://img.shields.io/github/license/qdm12/gluetun)](https://github.com/qdm12/gluetun/blob/master/LICENSE)
+13 -35
View File
@@ -2,7 +2,6 @@ package main
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"io/fs" "io/fs"
"net/http" "net/http"
@@ -43,7 +42,6 @@ import (
"github.com/qdm12/gluetun/internal/server" "github.com/qdm12/gluetun/internal/server"
"github.com/qdm12/gluetun/internal/shadowsocks" "github.com/qdm12/gluetun/internal/shadowsocks"
"github.com/qdm12/gluetun/internal/storage" "github.com/qdm12/gluetun/internal/storage"
"github.com/qdm12/gluetun/internal/tun"
updater "github.com/qdm12/gluetun/internal/updater/loop" updater "github.com/qdm12/gluetun/internal/updater/loop"
"github.com/qdm12/gluetun/internal/updater/resolver" "github.com/qdm12/gluetun/internal/updater/resolver"
"github.com/qdm12/gluetun/internal/updater/unzip" "github.com/qdm12/gluetun/internal/updater/unzip"
@@ -80,7 +78,6 @@ func main() {
logger := log.New(log.SetLevel(log.LevelInfo)) logger := log.New(log.SetLevel(log.LevelInfo))
args := os.Args args := os.Args
tun := tun.New()
netLinkDebugLogger := logger.New(log.SetComponent("netlink")) netLinkDebugLogger := logger.New(log.SetComponent("netlink"))
netLinker := netlink.New(netLinkDebugLogger) netLinker := netlink.New(netLinkDebugLogger)
cli := cli.New() cli := cli.New()
@@ -100,7 +97,7 @@ func main() {
errorCh := make(chan error) errorCh := make(chan error)
go func() { go func() {
errorCh <- _main(ctx, buildInfo, args, logger, reader, tun, netLinker, cmder, cli) errorCh <- _main(ctx, buildInfo, args, logger, reader, netLinker, cmder, cli)
}() }()
// Wait for OS signal or run error // Wait for OS signal or run error
@@ -145,7 +142,7 @@ func main() {
//nolint:gocognit,gocyclo,maintidx //nolint:gocognit,gocyclo,maintidx
func _main(ctx context.Context, buildInfo models.BuildInformation, func _main(ctx context.Context, buildInfo models.BuildInformation,
args []string, logger log.LoggerInterface, reader *reader.Reader, args []string, logger log.LoggerInterface, reader *reader.Reader,
tun Tun, netLinker netLinker, cmder RunStarter, netLinker netLinker, cmder RunStarter,
cli clier, cli clier,
) error { ) error {
if len(args) > 1 { // cli operation if len(args) > 1 { // cli operation
@@ -169,21 +166,19 @@ func _main(ctx context.Context, buildInfo models.BuildInformation,
defer fmt.Println(gluetunLogo) defer fmt.Println(gluetunLogo)
announcementExp, err := time.Parse(time.RFC3339, "2026-04-30T00:00:00Z") announcementExp, err := time.Parse(time.RFC3339, "2026-06-30T00:00:00Z")
if err != nil { if err != nil {
return err return err
} }
splashSettings := gosplash.Settings{ splashSettings := gosplash.Settings{
User: "qdm12", User: "qdm12",
Repository: "gluetun", Repository: "gluetun",
Emails: []string{"quentin.mcgaw@gmail.com"}, Emails: []string{"quentin.mcgaw@gmail.com"},
Version: buildInfo.Version, Version: buildInfo.Version,
Commit: buildInfo.Commit, Commit: buildInfo.Commit,
Created: buildInfo.Created, Created: buildInfo.Created,
Announcement: "the repository will be migrated to https://github.com/passteque/gluetun on 2026-05-21, " + Announcement: "Your servers data files are now migrated to /gluetun/servers/",
"which is a Github organization under my sole control, so don't get alarmed if you get redirected " + AnnounceExp: announcementExp,
"in the coming days (reason: personal paperwork ugh)",
AnnounceExp: announcementExp,
// Sponsor information // Sponsor information
PaypalUser: "qmcgaw", PaypalUser: "qmcgaw",
GithubSponsor: "qdm12", GithubSponsor: "qdm12",
@@ -245,7 +240,8 @@ func _main(ctx context.Context, buildInfo models.BuildInformation,
// TODO run this in a loop or in openvpn to reload from file without restarting // TODO run this in a loop or in openvpn to reload from file without restarting
storageLogger := logger.New(log.SetComponent("storage")) storageLogger := logger.New(log.SetComponent("storage"))
storage, err := storage.New(storageLogger, *allSettings.Storage.Filepath) storage, err := storage.New(storageLogger, *allSettings.Storage.ServersEnabled,
allSettings.Storage.ServersPath, allSettings.Storage.LegacyServersFilepath)
if err != nil { if err != nil {
return err return err
} }
@@ -343,19 +339,6 @@ func _main(ctx context.Context, buildInfo models.BuildInformation,
return fmt.Errorf("adding local rules: %w", err) return fmt.Errorf("adding local rules: %w", err)
} }
const tunDevice = "/dev/net/tun"
err = tun.Check(tunDevice)
if err != nil {
if !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("checking TUN device: %w (see the Wiki errors/tun page)", err)
}
logger.Info(err.Error() + "; creating it...")
err = tun.Create(tunDevice)
if err != nil {
return fmt.Errorf("creating tun device: %w", err)
}
}
for _, port := range allSettings.Firewall.InputPorts { for _, port := range allSettings.Firewall.InputPorts {
for _, defaultRoute := range defaultRoutes { for _, defaultRoute := range defaultRoutes {
err = firewallConf.SetAllowedPort(ctx, port, defaultRoute.NetInterface) err = firewallConf.SetAllowedPort(ctx, port, defaultRoute.NetInterface)
@@ -627,11 +610,6 @@ type clier interface {
GenKey(args []string) error GenKey(args []string) error
} }
type Tun interface {
Check(tunDevice string) error
Create(tunDevice string) error
}
type RunStarter interface { type RunStarter interface {
Run(cmd *exec.Cmd) (output string, err error) Run(cmd *exec.Cmd) (output string, err error)
Start(cmd *exec.Cmd) (stdoutLines, stderrLines <-chan string, Start(cmd *exec.Cmd) (stdoutLines, stderrLines <-chan string,
+2 -1
View File
@@ -15,6 +15,7 @@ require (
github.com/mdlayher/netlink v1.9.0 github.com/mdlayher/netlink v1.9.0
github.com/pelletier/go-toml/v2 v2.2.4 github.com/pelletier/go-toml/v2 v2.2.4
github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a
github.com/qdm12/gluetun-servers v0.1.0
github.com/qdm12/gosettings v0.4.4 github.com/qdm12/gosettings v0.4.4
github.com/qdm12/goshutdown v0.3.0 github.com/qdm12/goshutdown v0.3.0
github.com/qdm12/gosplash v0.2.1-0.20260305164749-b713de4fee6c github.com/qdm12/gosplash v0.2.1-0.20260305164749-b713de4fee6c
@@ -26,6 +27,7 @@ require (
github.com/ulikunitz/xz v0.5.15 github.com/ulikunitz/xz v0.5.15
github.com/youmark/pkcs8 v0.0.0-20201027041543-1326539a0a0a github.com/youmark/pkcs8 v0.0.0-20201027041543-1326539a0a0a
golang.org/x/exp v0.0.0-20241009180824-f66d83c29e7c golang.org/x/exp v0.0.0-20241009180824-f66d83c29e7c
golang.org/x/mod v0.33.0
golang.org/x/net v0.51.0 golang.org/x/net v0.51.0
golang.org/x/sys v0.42.0 golang.org/x/sys v0.42.0
golang.org/x/text v0.35.0 golang.org/x/text v0.35.0
@@ -57,7 +59,6 @@ require (
github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978 // indirect github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978 // indirect
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3 // indirect github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3 // indirect
golang.org/x/crypto v0.48.0 // indirect golang.org/x/crypto v0.48.0 // indirect
golang.org/x/mod v0.33.0 // indirect
golang.org/x/sync v0.20.0 // indirect golang.org/x/sync v0.20.0 // indirect
golang.org/x/tools v0.42.0 // indirect golang.org/x/tools v0.42.0 // indirect
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
+2
View File
@@ -76,6 +76,8 @@ github.com/prometheus/procfs v0.15.1 h1:YagwOFzUgYfKKHX6Dr+sHT7km/hxC76UB0leargg
github.com/prometheus/procfs v0.15.1/go.mod h1:fB45yRUv8NstnjriLhBQLuOUt+WW4BsoGhij/e3PBqk= github.com/prometheus/procfs v0.15.1/go.mod h1:fB45yRUv8NstnjriLhBQLuOUt+WW4BsoGhij/e3PBqk=
github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a h1:TE157yPQmAbVruH0MWCQzs0vTT/6t96DkoWUXd6PVuc= github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a h1:TE157yPQmAbVruH0MWCQzs0vTT/6t96DkoWUXd6PVuc=
github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a/go.mod h1:98foWgXJZ+g8gJIuO+fdO+oWpFei5WShMFTeN4Im2lE= github.com/qdm12/dns/v2 v2.0.0-rc9.0.20260421173011-9de8e7fdbe3a/go.mod h1:98foWgXJZ+g8gJIuO+fdO+oWpFei5WShMFTeN4Im2lE=
github.com/qdm12/gluetun-servers v0.1.0 h1:w9JLghKZwI0Gzpp9p5rNANgEYUUZ1dxdxsG6NKIojaY=
github.com/qdm12/gluetun-servers v0.1.0/go.mod h1:acttuyHyoFDu6GTbf3kAV+QXeiX8oJeh0MBic67/9z8=
github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978 h1:TRGpCU1l0lNwtogEUSs5U+RFceYxkAJUmrGabno7J5c= github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978 h1:TRGpCU1l0lNwtogEUSs5U+RFceYxkAJUmrGabno7J5c=
github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978/go.mod h1:D1Po4CRQLYjccnAR2JsVlN1sBMgQrcNLONbvyuzcdTg= github.com/qdm12/goservices v0.1.1-0.20251104135713-6bee97bd4978/go.mod h1:D1Po4CRQLYjccnAR2JsVlN1sBMgQrcNLONbvyuzcdTg=
github.com/qdm12/gosettings v0.4.4 h1:SM6tOZDf6k8qbjWU8KWyBF4mWIixfsKCfh9DGRLHlj4= github.com/qdm12/gosettings v0.4.4 h1:SM6tOZDf6k8qbjWU8KWyBF4mWIixfsKCfh9DGRLHlj4=
+1 -1
View File
@@ -31,7 +31,7 @@ type urlData struct{}
func New(client *http.Client, logger Logger, settings settings.BoringPoll) *BoringPoll { func New(client *http.Client, logger Logger, settings settings.BoringPoll) *BoringPoll {
urlToData := make(map[string]*urlData) urlToData := make(map[string]*urlData)
if *settings.GluetunCom { if *settings.GluetunCom {
urlToData["https://gluetun.com/wp-json"] = &urlData{} logger.Infof("gluetun.com is DOWN most likely thanks to you! so not doing anything anymore")
} }
return &BoringPoll{ return &BoringPoll{
client: client, client: client,
+2 -6
View File
@@ -1,11 +1,7 @@
package cli package cli
type CLI struct { type CLI struct{}
repoServersPath string
}
func New() *CLI { func New() *CLI {
return &CLI{ return &CLI{}
repoServersPath: "./internal/storage/servers.json",
}
} }
+2 -5
View File
@@ -9,9 +9,7 @@ import (
"path/filepath" "path/filepath"
"strings" "strings"
"github.com/qdm12/gluetun/internal/constants"
"github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/constants/providers"
"github.com/qdm12/gluetun/internal/storage"
"golang.org/x/text/cases" "golang.org/x/text/cases"
"golang.org/x/text/language" "golang.org/x/text/language"
) )
@@ -74,10 +72,9 @@ func (c *CLI) FormatServers(args []string) error {
} }
} }
logger := newNoopLogger() storage, err := setupStorage(newNoopLogger())
storage, err := storage.New(logger, constants.ServersData)
if err != nil { if err != nil {
return fmt.Errorf("creating servers storage: %w", err) return fmt.Errorf("setting up storage: %w", err)
} }
formatted, err := storage.Format(providerToFormat, format) formatted, err := storage.Format(providerToFormat, format)
+39
View File
@@ -0,0 +1,39 @@
package cli
import (
"fmt"
"github.com/qdm12/gluetun/internal/configuration/settings"
"github.com/qdm12/gluetun/internal/configuration/sources/files"
"github.com/qdm12/gluetun/internal/configuration/sources/secrets"
"github.com/qdm12/gluetun/internal/storage"
"github.com/qdm12/gosettings/reader"
"github.com/qdm12/gosettings/reader/sources/env"
)
type storageSetupLogger interface {
storage.Logger
files.Warner
}
func setupStorage(logger storageSetupLogger) (s *storage.Storage, err error) {
settingsReader := reader.New(reader.Settings{
Sources: []reader.Source{
secrets.New(logger),
files.New(logger),
env.New(env.Settings{}),
},
})
var settings settings.Storage
err = settings.Read(settingsReader)
if err != nil {
return nil, fmt.Errorf("reading storage settings: %w", err)
}
settings.SetDefaults()
storage, err := storage.New(logger, *settings.ServersEnabled, settings.ServersPath,
settings.LegacyServersFilepath)
if err != nil {
return nil, fmt.Errorf("creating storage: %w", err)
}
return storage, nil
}
+4 -2
View File
@@ -6,5 +6,7 @@ func newNoopLogger() *noopLogger {
return new(noopLogger) return new(noopLogger)
} }
func (l *noopLogger) Info(string) {} func (l *noopLogger) Info(string) {}
func (l *noopLogger) Warn(string) {} func (l *noopLogger) Infof(string, ...any) {}
func (l *noopLogger) Warn(string) {}
func (l *noopLogger) Warnf(string, ...any) {}
+2 -4
View File
@@ -9,12 +9,10 @@ import (
"time" "time"
"github.com/qdm12/gluetun/internal/configuration/settings" "github.com/qdm12/gluetun/internal/configuration/settings"
"github.com/qdm12/gluetun/internal/constants"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
"github.com/qdm12/gluetun/internal/netlink" "github.com/qdm12/gluetun/internal/netlink"
"github.com/qdm12/gluetun/internal/openvpn/extract" "github.com/qdm12/gluetun/internal/openvpn/extract"
"github.com/qdm12/gluetun/internal/provider" "github.com/qdm12/gluetun/internal/provider"
"github.com/qdm12/gluetun/internal/storage"
"github.com/qdm12/gluetun/internal/updater/resolver" "github.com/qdm12/gluetun/internal/updater/resolver"
"github.com/qdm12/gosettings/reader" "github.com/qdm12/gosettings/reader"
) )
@@ -49,9 +47,9 @@ type IPv6Checker interface {
func (c *CLI) OpenvpnConfig(logger OpenvpnConfigLogger, reader *reader.Reader, func (c *CLI) OpenvpnConfig(logger OpenvpnConfigLogger, reader *reader.Reader,
ipv6Checker IPv6Checker, ipv6Checker IPv6Checker,
) error { ) error {
storage, err := storage.New(logger, constants.ServersData) storage, err := setupStorage(newNoopLogger())
if err != nil { if err != nil {
return err return fmt.Errorf("setting up storage: %w", err)
} }
var allSettings settings.Settings var allSettings settings.Settings
+15 -24
View File
@@ -13,12 +13,10 @@ import (
"github.com/qdm12/dns/v2/pkg/doh" "github.com/qdm12/dns/v2/pkg/doh"
dnsprovider "github.com/qdm12/dns/v2/pkg/provider" dnsprovider "github.com/qdm12/dns/v2/pkg/provider"
"github.com/qdm12/gluetun/internal/configuration/settings" "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/providers"
"github.com/qdm12/gluetun/internal/openvpn/extract" "github.com/qdm12/gluetun/internal/openvpn/extract"
"github.com/qdm12/gluetun/internal/provider" "github.com/qdm12/gluetun/internal/provider"
"github.com/qdm12/gluetun/internal/publicip/api" "github.com/qdm12/gluetun/internal/publicip/api"
"github.com/qdm12/gluetun/internal/storage"
"github.com/qdm12/gluetun/internal/updater" "github.com/qdm12/gluetun/internal/updater"
"github.com/qdm12/gluetun/internal/updater/resolver" "github.com/qdm12/gluetun/internal/updater/resolver"
"github.com/qdm12/gluetun/internal/updater/unzip" "github.com/qdm12/gluetun/internal/updater/unzip"
@@ -26,18 +24,19 @@ import (
type UpdaterLogger interface { type UpdaterLogger interface {
Info(s string) Info(s string)
Infof(format string, args ...any)
Warn(s string) Warn(s string)
Warnf(format string, args ...any)
Error(s string) Error(s string)
} }
func (c *CLI) Update(ctx context.Context, args []string, logger UpdaterLogger) error { func (c *CLI) Update(ctx context.Context, args []string, logger UpdaterLogger) error {
options := settings.Updater{} options := settings.Updater{}
var endUserMode, maintainerMode, updateAll bool // TODO v4: remove flags below already present in standard settings
var endUserMode, maintainerMode bool
var updateAll bool
var dnsServer, csvProviders, ipToken, protonUsername, protonEmail, protonPassword string var dnsServer, csvProviders, ipToken, protonUsername, protonEmail, protonPassword string
flagSet := flag.NewFlagSet("update", flag.ExitOnError) flagSet := flag.NewFlagSet("update", flag.ExitOnError)
flagSet.BoolVar(&endUserMode, "enduser", false, "Write results to /gluetun/servers.json (for end users)")
flagSet.BoolVar(&maintainerMode, "maintainer", false,
"Write results to ./internal/storage/servers.json to modify the program (for maintainers)")
flagSet.StringVar(&dnsServer, "dns", "", "no longer used, your DNS will use DoH with Cloudflare and Google") flagSet.StringVar(&dnsServer, "dns", "", "no longer used, your DNS will use DoH with Cloudflare and Google")
const defaultMinRatio = 0.8 const defaultMinRatio = 0.8
flagSet.Float64Var(&options.MinRatio, "minratio", defaultMinRatio, flagSet.Float64Var(&options.MinRatio, "minratio", defaultMinRatio,
@@ -49,16 +48,19 @@ func (c *CLI) Update(ctx context.Context, args []string, logger UpdaterLogger) e
"(Retro-compatibility) Username to use to authenticate with Proton. Use -proton-email instead.") // v4 remove this "(Retro-compatibility) Username to use to authenticate with Proton. Use -proton-email instead.") // v4 remove this
flagSet.StringVar(&protonEmail, "proton-email", "", "Email to use to authenticate with Proton") flagSet.StringVar(&protonEmail, "proton-email", "", "Email to use to authenticate with Proton")
flagSet.StringVar(&protonPassword, "proton-password", "", "Password to use to authenticate with Proton") flagSet.StringVar(&protonPassword, "proton-password", "", "Password to use to authenticate with Proton")
flagSet.BoolVar(&endUserMode, "enduser", false, "deprecated")
flagSet.BoolVar(&maintainerMode, "maintainer", false, "deprecated")
if err := flagSet.Parse(args); err != nil { if err := flagSet.Parse(args); err != nil {
return err return err
} }
if dnsServer != "" { switch {
case dnsServer != "":
logger.Warn("The -dns flag is no longer used, your DNS will use DoH with Cloudflare and Google") logger.Warn("The -dns flag is no longer used, your DNS will use DoH with Cloudflare and Google")
} case endUserMode:
logger.Warn("The -enduser flag is now unused")
if !endUserMode && !maintainerMode { case maintainerMode:
return errors.New("at least one of -enduser or -maintainer must be specified") logger.Warn("The -maintainer flag is now unused")
} }
if updateAll { if updateAll {
@@ -87,11 +89,7 @@ func (c *CLI) Update(ctx context.Context, args []string, logger UpdaterLogger) e
return fmt.Errorf("options validation failed: %w", err) return fmt.Errorf("options validation failed: %w", err)
} }
serversDataPath := constants.ServersData storage, err := setupStorage(logger)
if maintainerMode {
serversDataPath = ""
}
storage, err := storage.New(logger, serversDataPath)
if err != nil { if err != nil {
return fmt.Errorf("creating servers storage: %w", err) return fmt.Errorf("creating servers storage: %w", err)
} }
@@ -127,18 +125,11 @@ func (c *CLI) Update(ctx context.Context, args []string, logger UpdaterLogger) e
providers := provider.NewProviders(storage, time.Now, logger, httpClient, providers := provider.NewProviders(storage, time.Now, logger, httpClient,
unzipper, parallelResolver, ipFetcher, openvpnFileExtractor, options) unzipper, parallelResolver, ipFetcher, openvpnFileExtractor, options)
updater := updater.New(httpClient, storage, providers, logger) updater := updater.New(httpClient, storage, providers, logger, *options.PreferDirectDownload)
err = updater.UpdateServers(ctx, options.Providers, options.MinRatio) err = updater.UpdateServers(ctx, options.Providers, options.MinRatio)
if err != nil { if err != nil {
return fmt.Errorf("updating server information: %w", err) return fmt.Errorf("updating server information: %w", err)
} }
if maintainerMode {
err := storage.FlushToFile(c.repoServersPath)
if err != nil {
return fmt.Errorf("writing servers data to embedded JSON file: %w", err)
}
}
return nil return nil
} }
+9 -18
View File
@@ -21,7 +21,6 @@ func (c *Cmder) Start(cmd *exec.Cmd) (
func start(cmd execCmd) (stdoutLines, stderrLines <-chan string, func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
waitError <-chan error, startErr error, waitError <-chan error, startErr error,
) { ) {
stop := make(chan struct{})
stdoutReady := make(chan struct{}) stdoutReady := make(chan struct{})
stdoutLinesCh := make(chan string) stdoutLinesCh := make(chan string)
stdoutDone := make(chan struct{}) stdoutDone := make(chan struct{})
@@ -33,22 +32,20 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
if err != nil { if err != nil {
return nil, nil, nil, err return nil, nil, nil, err
} }
go streamToChannel(stdoutReady, stop, stdoutDone, stdout, stdoutLinesCh) go streamToChannel(stdoutReady, stdoutDone, stdout, stdoutLinesCh)
stderr, err := cmd.StderrPipe() stderr, err := cmd.StderrPipe()
if err != nil { if err != nil {
_ = stdout.Close() _ = stdout.Close()
close(stop)
<-stdoutDone <-stdoutDone
return nil, nil, nil, err return nil, nil, nil, err
} }
go streamToChannel(stderrReady, stop, stderrDone, stderr, stderrLinesCh) go streamToChannel(stderrReady, stderrDone, stderr, stderrLinesCh)
err = cmd.Start() err = cmd.Start()
if err != nil { if err != nil {
_ = stdout.Close() _ = stdout.Close()
_ = stderr.Close() _ = stderr.Close()
close(stop)
<-stdoutDone <-stdoutDone
<-stderrDone <-stderrDone
return nil, nil, nil, err return nil, nil, nil, err
@@ -57,19 +54,20 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
waitErrorCh := make(chan error) waitErrorCh := make(chan error)
go func() { go func() {
err := cmd.Wait() err := cmd.Wait()
_ = stdout.Close()
_ = stderr.Close()
close(stop)
<-stdoutDone <-stdoutDone
<-stderrDone <-stderrDone
_ = stdout.Close()
_ = stderr.Close()
waitErrorCh <- err waitErrorCh <- err
}() }()
<-stdoutReady
<-stderrReady
return stdoutLinesCh, stderrLinesCh, waitErrorCh, nil return stdoutLinesCh, stderrLinesCh, waitErrorCh, nil
} }
func streamToChannel(ready chan<- struct{}, func streamToChannel(ready chan<- struct{}, done chan<- struct{},
stop <-chan struct{}, done chan<- struct{},
stream io.Reader, lines chan<- string, stream io.Reader, lines chan<- string,
) { ) {
defer close(done) defer close(done)
@@ -89,12 +87,5 @@ func streamToChannel(ready chan<- struct{},
if err == nil || errors.Is(err, os.ErrClosed) { if err == nil || errors.Is(err, os.ErrClosed) {
return return
} }
lines <- "stream error: " + err.Error()
// ignore the error if it is stopped.
select {
case <-stop:
return
default:
lines <- "stream error: " + err.Error()
}
} }
+2 -2
View File
@@ -132,7 +132,7 @@ func (s *Settings) SetDefaults() {
s.IPv6.setDefaults() s.IPv6.setDefaults()
s.PublicIP.setDefaults() s.PublicIP.setDefaults()
s.Shadowsocks.setDefaults() s.Shadowsocks.setDefaults()
s.Storage.setDefaults() s.Storage.SetDefaults()
s.System.setDefaults() s.System.setDefaults()
s.Version.setDefaults() s.Version.setDefaults()
s.VPN.setDefaults() s.VPN.setDefaults()
@@ -213,7 +213,7 @@ func (s *Settings) Read(r *reader.Reader, warner Warner) (err error) {
return s.PublicIP.read(r, warner) return s.PublicIP.read(r, warner)
}, },
"shadowsocks": s.Shadowsocks.read, "shadowsocks": s.Shadowsocks.read,
"storage": s.Storage.read, "storage": s.Storage.Read,
"system": s.System.read, "system": s.System.read,
"updater": s.Updater.read, "updater": s.Updater.read,
"version": s.Version.read, "version": s.Version.read,
@@ -90,7 +90,7 @@ func Test_Settings_String(t *testing.T) {
| ├── Logging: yes | ├── Logging: yes
| └── Authentication file path: /gluetun/auth/config.toml | └── Authentication file path: /gluetun/auth/config.toml
├── Storage settings: ├── Storage settings:
| └── Filepath: /gluetun/servers.json | └── Servers directory path: /gluetun/servers/
├── OS Alpine settings: ├── OS Alpine settings:
| ├── Process UID: 1000 | ├── Process UID: 1000
| └── Process GID: 1000 | └── Process GID: 1000
+51 -14
View File
@@ -11,15 +11,26 @@ import (
// Storage contains settings to configure the storage. // Storage contains settings to configure the storage.
type Storage struct { type Storage struct {
// Filepath is the path to the servers.json file. An empty string disables on-disk storage. // ServersEnabled is whether to enable storage of servers on disk.
Filepath *string // It defaults to true.
ServersEnabled *bool
// ServersPath is the path to the servers files directory, and cannot be
// the empty string.
ServersPath string
// LegacyServersFilepath is the legacy "fat" JSON filepath to migrate from.
// TODO v4: remove
LegacyServersFilepath string
} }
func (s Storage) validate() (err error) { func (s Storage) validate() (err error) {
if *s.Filepath != "" { // optional if *s.ServersEnabled {
_, err := filepath.Abs(*s.Filepath) _, err := filepath.Abs(s.ServersPath)
if err != nil { if err != nil {
return fmt.Errorf("filepath is not valid: %w", err) return fmt.Errorf("servers path is not valid: %w", err)
}
_, err = filepath.Abs(s.LegacyServersFilepath)
if err != nil {
return fmt.Errorf("legacy servers filepath is not valid: %w", err)
} }
} }
return nil return nil
@@ -27,17 +38,25 @@ func (s Storage) validate() (err error) {
func (s *Storage) copy() (copied Storage) { func (s *Storage) copy() (copied Storage) {
return Storage{ return Storage{
Filepath: gosettings.CopyPointer(s.Filepath), ServersEnabled: gosettings.CopyPointer(s.ServersEnabled),
ServersPath: s.ServersPath,
LegacyServersFilepath: s.LegacyServersFilepath,
} }
} }
func (s *Storage) overrideWith(other Storage) { func (s *Storage) overrideWith(other Storage) {
s.Filepath = gosettings.OverrideWithPointer(s.Filepath, other.Filepath) s.ServersEnabled = gosettings.OverrideWithPointer(s.ServersEnabled, other.ServersEnabled)
s.ServersPath = gosettings.OverrideWithComparable(s.ServersPath, other.ServersPath)
s.LegacyServersFilepath = gosettings.OverrideWithComparable(s.LegacyServersFilepath, other.LegacyServersFilepath)
} }
func (s *Storage) setDefaults() { const defaultLegacyServersFilepath = "/gluetun/servers.json"
const defaultFilepath = "/gluetun/servers.json"
s.Filepath = gosettings.DefaultPointer(s.Filepath, defaultFilepath) func (s *Storage) SetDefaults() {
s.ServersEnabled = gosettings.DefaultPointer(s.ServersEnabled, true)
const defaultServersPath = "/gluetun/servers/"
s.ServersPath = gosettings.DefaultComparable(s.ServersPath, defaultServersPath)
s.LegacyServersFilepath = gosettings.DefaultComparable(s.LegacyServersFilepath, defaultLegacyServersFilepath)
} }
func (s Storage) String() string { func (s Storage) String() string {
@@ -45,15 +64,33 @@ func (s Storage) String() string {
} }
func (s Storage) toLinesNode() (node *gotree.Node) { func (s Storage) toLinesNode() (node *gotree.Node) {
if *s.Filepath == "" { if !*s.ServersEnabled {
return gotree.New("Storage settings: disabled") return gotree.New("Storage settings: disabled")
} }
node = gotree.New("Storage settings:") node = gotree.New("Storage settings:")
node.Appendf("Filepath: %s", *s.Filepath) node.Appendf("Servers directory path: %s", s.ServersPath)
if s.LegacyServersFilepath != defaultLegacyServersFilepath {
node.Appendf("Legacy servers filepath: %s", s.LegacyServersFilepath)
}
return node return node
} }
func (s *Storage) read(r *reader.Reader) (err error) { func (s *Storage) Read(r *reader.Reader) (err error) {
s.Filepath = r.Get("STORAGE_FILEPATH", reader.AcceptEmpty(true)) // Retro-compatibility:
// TODO v4: remove support for STORAGE_FILEPATH
filePath := r.Get("STORAGE_FILEPATH", reader.AcceptEmpty(true), reader.IsRetro("STORAGE_SERVERS_DIRECTORY_PATH"))
if filePath != nil {
if *filePath == "" {
s.ServersEnabled = ptrTo(false)
} else {
s.LegacyServersFilepath = *filePath
}
} else {
s.ServersEnabled, err = r.BoolPtr("STORAGE_SERVERS_ENABLED")
if err != nil {
return err
}
s.ServersPath = r.String("STORAGE_SERVERS_DIRECTORY_PATH")
}
return nil return nil
} }
+17 -5
View File
@@ -29,6 +29,9 @@ type Updater struct {
// Providers is the list of VPN service providers // Providers is the list of VPN service providers
// to update server information for. // to update server information for.
Providers []string Providers []string
// PreferDirectDownload is whether to prefer direct download of
// server data from Github (recommended).
PreferDirectDownload *bool
// ProtonEmail is the email to authenticate with the Proton API. // ProtonEmail is the email to authenticate with the Proton API.
ProtonEmail *string ProtonEmail *string
// ProtonPassword is the password to authenticate with the Proton API. // ProtonPassword is the password to authenticate with the Proton API.
@@ -72,11 +75,12 @@ func (u Updater) Validate() (err error) {
func (u *Updater) copy() (copied Updater) { func (u *Updater) copy() (copied Updater) {
return Updater{ return Updater{
Period: gosettings.CopyPointer(u.Period), Period: gosettings.CopyPointer(u.Period),
MinRatio: u.MinRatio, MinRatio: u.MinRatio,
Providers: gosettings.CopySlice(u.Providers), Providers: gosettings.CopySlice(u.Providers),
ProtonEmail: gosettings.CopyPointer(u.ProtonEmail), PreferDirectDownload: gosettings.CopyPointer(u.PreferDirectDownload),
ProtonPassword: gosettings.CopyPointer(u.ProtonPassword), ProtonEmail: gosettings.CopyPointer(u.ProtonEmail),
ProtonPassword: gosettings.CopyPointer(u.ProtonPassword),
} }
} }
@@ -87,6 +91,7 @@ func (u *Updater) overrideWith(other Updater) {
u.Period = gosettings.OverrideWithPointer(u.Period, other.Period) u.Period = gosettings.OverrideWithPointer(u.Period, other.Period)
u.MinRatio = gosettings.OverrideWithComparable(u.MinRatio, other.MinRatio) u.MinRatio = gosettings.OverrideWithComparable(u.MinRatio, other.MinRatio)
u.Providers = gosettings.OverrideWithSlice(u.Providers, other.Providers) u.Providers = gosettings.OverrideWithSlice(u.Providers, other.Providers)
u.PreferDirectDownload = gosettings.OverrideWithPointer(u.PreferDirectDownload, other.PreferDirectDownload)
u.ProtonEmail = gosettings.OverrideWithPointer(u.ProtonEmail, other.ProtonEmail) u.ProtonEmail = gosettings.OverrideWithPointer(u.ProtonEmail, other.ProtonEmail)
u.ProtonPassword = gosettings.OverrideWithPointer(u.ProtonPassword, other.ProtonPassword) u.ProtonPassword = gosettings.OverrideWithPointer(u.ProtonPassword, other.ProtonPassword)
} }
@@ -104,6 +109,7 @@ func (u *Updater) SetDefaults(vpnProvider string) {
} }
// Set these to empty strings to avoid nil pointer panics // Set these to empty strings to avoid nil pointer panics
u.PreferDirectDownload = gosettings.DefaultPointer(u.PreferDirectDownload, false)
u.ProtonEmail = gosettings.DefaultPointer(u.ProtonEmail, "") u.ProtonEmail = gosettings.DefaultPointer(u.ProtonEmail, "")
u.ProtonPassword = gosettings.DefaultPointer(u.ProtonPassword, "") u.ProtonPassword = gosettings.DefaultPointer(u.ProtonPassword, "")
} }
@@ -121,6 +127,7 @@ func (u Updater) toLinesNode() (node *gotree.Node) {
node.Appendf("Update period: %s", *u.Period) node.Appendf("Update period: %s", *u.Period)
node.Appendf("Minimum ratio: %.1f", u.MinRatio) node.Appendf("Minimum ratio: %.1f", u.MinRatio)
node.Appendf("Providers to update: %s", strings.Join(u.Providers, ", ")) node.Appendf("Providers to update: %s", strings.Join(u.Providers, ", "))
node.Appendf("Prefer direct download: %s", gosettings.BoolToYesNo(u.PreferDirectDownload))
if slices.Contains(u.Providers, providers.Protonvpn) { if slices.Contains(u.Providers, providers.Protonvpn) {
node.Appendf("Proton API email: %s", *u.ProtonEmail) node.Appendf("Proton API email: %s", *u.ProtonEmail)
node.Appendf("Proton API password: %s", gosettings.ObfuscateKey(*u.ProtonPassword)) node.Appendf("Proton API password: %s", gosettings.ObfuscateKey(*u.ProtonPassword))
@@ -142,6 +149,11 @@ func (u *Updater) read(r *reader.Reader) (err error) {
u.Providers = r.CSV("UPDATER_VPN_SERVICE_PROVIDERS") u.Providers = r.CSV("UPDATER_VPN_SERVICE_PROVIDERS")
u.PreferDirectDownload, err = r.BoolPtr("UPDATER_PREFER_DIRECT_DOWNLOAD")
if err != nil {
return err
}
u.ProtonEmail = r.Get("UPDATER_PROTONVPN_EMAIL") u.ProtonEmail = r.Get("UPDATER_PROTONVPN_EMAIL")
if u.ProtonEmail == nil { if u.ProtonEmail == nil {
protonUsername := r.String("UPDATER_PROTONVPN_USERNAME", reader.IsRetro("UPDATER_PROTONVPN_EMAIL")) protonUsername := r.String("UPDATER_PROTONVPN_USERNAME", reader.IsRetro("UPDATER_PROTONVPN_EMAIL"))
-6
View File
@@ -1,6 +0,0 @@
package constants
const (
// ServersData is the server information filepath.
ServersData = "/gluetun/servers.json"
)
+39 -5
View File
@@ -1,6 +1,7 @@
package iptables package iptables
import ( import (
"bufio"
"context" "context"
"fmt" "fmt"
"os/exec" "os/exec"
@@ -41,8 +42,7 @@ func (c *Config) saveAndRestore(ctx context.Context) (restore func(context.Conte
// Callers of saveAndRestoreIPv4 MUST always lock the [Config] iptablesMutex // Callers of saveAndRestoreIPv4 MUST always lock the [Config] iptablesMutex
// before calling this function. // before calling this function.
func (c *Config) saveAndRestoreIPv4(ctx context.Context) (restore func(context.Context), err error) { func (c *Config) saveAndRestoreIPv4(ctx context.Context) (restore func(context.Context), err error) {
cmd := exec.CommandContext(ctx, c.ipTables+"-save") //nolint:gosec data, err := saveData(ctx, c.ipTables)
data, err := c.runner.Run(cmd)
if err != nil { if err != nil {
return nil, fmt.Errorf("saving IPv4 iptables: %w", err) return nil, fmt.Errorf("saving IPv4 iptables: %w", err)
} }
@@ -65,14 +65,13 @@ func (c *Config) saveAndRestoreIPv6(ctx context.Context) (restore func(context.C
return nil, nil //nolint:nilnil return nil, nil //nolint:nilnil
} }
cmd := exec.CommandContext(ctx, c.ip6Tables+"-save") //nolint:gosec data, err := saveData(ctx, c.ip6Tables)
data, err := c.runner.Run(cmd)
if err != nil { if err != nil {
return nil, fmt.Errorf("saving IPv6 iptables: %w", err) return nil, fmt.Errorf("saving IPv6 iptables: %w", err)
} }
restore = func(ctx context.Context) { restore = func(ctx context.Context) {
cmd = exec.CommandContext(ctx, c.ip6Tables+"-restore") //nolint:gosec cmd := exec.CommandContext(ctx, c.ip6Tables+"-restore") //nolint:gosec
cmd.Stdin = strings.NewReader(data) cmd.Stdin = strings.NewReader(data)
output, err := c.runner.Run(cmd) output, err := c.runner.Run(cmd)
if err != nil { if err != nil {
@@ -85,3 +84,38 @@ func (c *Config) saveAndRestoreIPv6(ctx context.Context) (restore func(context.C
func makeRestoreErrorMessage(err error, output, data string) string { func makeRestoreErrorMessage(err error, output, data string) string {
return fmt.Sprintf("%s: %s: restoring from data:\n%s", err, output, data) return fmt.Sprintf("%s: %s: restoring from data:\n%s", err, output, data)
} }
func saveData(ctx context.Context, binary string) (data string, err error) {
cmd := exec.CommandContext(ctx, binary+"-save") //nolint:gosec
output, err := cmd.Output()
if err != nil {
if exitErr, ok := err.(*exec.ExitError); ok {
stderr := strings.TrimSuffix(string(exitErr.Stderr), "\n")
if stderr != "" {
return "", fmt.Errorf("running %s-save: %w: %s", binary, err, stderr)
}
}
return "", fmt.Errorf("running %s-save: %w", binary, err)
}
err = checkData(string(output))
if err != nil {
return "", fmt.Errorf("checking saved data: %w", err)
}
return string(output), nil
}
func checkData(data string) error {
scanner := bufio.NewScanner(strings.NewReader(data))
i := 0
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, "[unsupported") {
return fmt.Errorf("unsupported revision marker found in line %d: %s", i+1, line)
}
i++
}
if scanner.Err() != nil {
return fmt.Errorf("scanning data: %w", scanner.Err())
}
return nil
}
+2
View File
@@ -154,6 +154,8 @@ func (a *AllServers) Count() (count int) {
type Servers struct { type Servers struct {
Version uint16 `json:"version"` Version uint16 `json:"version"`
Timestamp int64 `json:"timestamp"` Timestamp int64 `json:"timestamp"`
Preferred bool `json:"preferred,omitempty"`
Filepath string `json:"filepath,omitempty"`
Servers []Server `json:"servers,omitempty"` Servers []Server `json:"servers,omitempty"`
} }
+2
View File
@@ -53,6 +53,8 @@ func extractDataFromLines(lines []string) (
func extractDataFromLine(line string) ( func extractDataFromLine(line string) (
ip netip.Addr, port uint16, protocol string, err error, ip netip.Addr, port uint16, protocol string, err error,
) { ) {
line = strings.TrimSpace(line)
switch { switch {
case strings.HasPrefix(line, "proto "): case strings.HasPrefix(line, "proto "):
protocol, err = extractProto(line) protocol, err = extractProto(line)
+8
View File
@@ -62,6 +62,14 @@ func Test_extractDataFromLines(t *testing.T) {
Protocol: constants.UDP, Protocol: constants.UDP,
}, },
}, },
"leading_whitespace": {
lines: []string{" proto tcp", "\tremote 1.2.3.4 443 tcp"},
connection: models.Connection{
IP: netip.AddrFrom4([4]byte{1, 2, 3, 4}),
Port: 443,
Protocol: constants.TCP,
},
},
} }
for name, testCase := range testCases { for name, testCase := range testCases {
+1 -1
View File
@@ -46,7 +46,7 @@ Your credentials might be wrong 🤨
` `
level = levelError level = levelError
case strings.Contains(s, "TLS Error: TLS key negotiation failed to occur within 60 seconds (check your network connectivity)"): //nolint:lll case strings.Contains(s, "TLS Error: TLS key negotiation failed to occur within 20 seconds (check your network connectivity)"): //nolint:lll
filtered = s + ` filtered = s + `
🚒🚒🚒🚒🚒🚨🚨🚨🚨🚨🚨🚒🚒🚒🚒🚒 🚒🚒🚒🚒🚒🚨🚨🚨🚨🚨🚨🚒🚒🚒🚒🚒
That error usually happens because either: That error usually happens because either:
+2 -2
View File
@@ -52,9 +52,9 @@ func Test_processLogLine(t *testing.T) {
}, },
"TLS key negotiation error": { "TLS key negotiation error": {
s: "TLS Error: TLS key negotiation failed to occur within " + s: "TLS Error: TLS key negotiation failed to occur within " +
"60 seconds (check your network connectivity)", "20 seconds (check your network connectivity)",
filtered: "TLS Error: TLS key negotiation failed to occur within " + filtered: "TLS Error: TLS key negotiation failed to occur within " +
"60 seconds (check your network connectivity)" + ` "20 seconds (check your network connectivity)" + `
🚒🚒🚒🚒🚒🚨🚨🚨🚨🚨🚨🚒🚒🚒🚒🚒 🚒🚒🚒🚒🚒🚨🚨🚨🚨🚨🚨🚒🚒🚒🚒🚒
That error usually happens because either: That error usually happens because either:
+1 -8
View File
@@ -4,7 +4,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"net/netip" "net/netip"
"strings"
"github.com/jsimonetti/rtnetlink" "github.com/jsimonetti/rtnetlink"
"github.com/qdm12/gluetun/internal/pmtud/constants" "github.com/qdm12/gluetun/internal/pmtud/constants"
@@ -28,10 +27,7 @@ func SrcAddr(dst netip.AddrPort, proto int) (src netip.AddrPort, cleanup func(),
return netip.AddrPortFrom(srcAddr, srcPort), cleanup, nil return netip.AddrPortFrom(srcAddr, srcPort), cleanup, nil
} }
var ( var errNoRoute = errors.New("no route to destination")
errNoRoute = errors.New("no route to destination")
ErrNetworkUnreachable = errors.New("network unreachable")
)
func srcIP(dst netip.Addr) (netip.Addr, error) { func srcIP(dst netip.Addr) (netip.Addr, error) {
conn, err := rtnetlink.Dial(nil) conn, err := rtnetlink.Dial(nil)
@@ -54,9 +50,6 @@ func srcIP(dst netip.Addr) (netip.Addr, error) {
} }
messages, err := conn.Route.Get(requestMessage) messages, err := conn.Route.Get(requestMessage)
if err != nil { if err != nil {
if strings.Contains(err.Error(), "network is unreachable") {
err = ErrNetworkUnreachable
}
return netip.Addr{}, fmt.Errorf("getting routes to %s: %w", dst, err) return netip.Addr{}, fmt.Errorf("getting routes to %s: %w", dst, err)
} }
-3
View File
@@ -43,9 +43,6 @@ func findHighestMSSDestination(ctx context.Context, familyToFD map[int]fileDescr
case err != nil: // error already occurred for another findMSS goroutine case err != nil: // error already occurred for another findMSS goroutine
case errors.Is(result.err, iptables.ErrMarkMatchModuleMissing): case errors.Is(result.err, iptables.ErrMarkMatchModuleMissing):
err = fmt.Errorf("finding MSS for %s: %w", result.dst, result.err) err = fmt.Errorf("finding MSS for %s: %w", result.dst, result.err)
case dst.Addr().Is6() && errors.Is(result.err, ip.ErrNetworkUnreachable):
// silently discard IPv6 network unreachable errors since they are common
// and expected when the host doesn't have IPv6 connectivity
default: // another error not due to the match module missing default: // another error not due to the match module missing
logger.Debugf("finding MSS for %s failed: %s", result.dst, result.err) logger.Debugf("finding MSS for %s failed: %s", result.dst, result.err)
} }
+3 -5
View File
@@ -1,21 +1,19 @@
package pmtud package pmtud
import ( import (
"net/netip"
"github.com/qdm12/gluetun/internal/constants" "github.com/qdm12/gluetun/internal/constants"
"github.com/qdm12/gluetun/internal/constants/vpn" "github.com/qdm12/gluetun/internal/constants/vpn"
pconstants "github.com/qdm12/gluetun/internal/pmtud/constants" pconstants "github.com/qdm12/gluetun/internal/pmtud/constants"
) )
// MaxTheoreticalVPNMTU returns the theoretical maximum MTU for a VPN tunnel // MaxTheoreticalVPNMTU returns the theoretical maximum MTU for a VPN tunnel
// given the VPN type, network protocol, and VPN gateway IP address. // given the VPN type, network protocol, and whether IPv6 is used.
// This is notably useful to skip testing MTU values higher than this value. // This is notably useful to skip testing MTU values higher than this value.
// The function panics if the network or VPN type is unknown. // The function panics if the network or VPN type is unknown.
func MaxTheoreticalVPNMTU(vpnType, network string, vpnGateway netip.Addr) uint32 { func MaxTheoreticalVPNMTU(vpnType, network string, ipv6 bool) uint32 {
const physicalLinkMTU = pconstants.MaxEthernetFrameSize const physicalLinkMTU = pconstants.MaxEthernetFrameSize
vpnLinkMTU := physicalLinkMTU vpnLinkMTU := physicalLinkMTU
if vpnGateway.Is4() { if !ipv6 {
vpnLinkMTU -= pconstants.IPv4HeaderLength vpnLinkMTU -= pconstants.IPv4HeaderLength
} else { } else {
vpnLinkMTU -= pconstants.IPv6HeaderLength vpnLinkMTU -= pconstants.IPv6HeaderLength
+29
View File
@@ -2,6 +2,10 @@ package service
import ( import (
"fmt" "fmt"
"maps"
"slices"
"sort"
"strconv"
"strings" "strings"
) )
@@ -20,3 +24,28 @@ func portsToString(ports []uint16) (s string) {
" and " + portStrings[len(portStrings)-1] " and " + portStrings[len(portStrings)-1]
} }
} }
func portPairsToString(internalToExternalPort map[uint16]uint16) (s string) {
switch len(internalToExternalPort) {
case 0:
return "no port forwarded"
case 1:
internal := slices.Collect(maps.Keys(internalToExternalPort))[0]
return "port forwarded is " + portPairToString(internal, internalToExternalPort[internal])
default:
portStrings := make([]string, 0, len(internalToExternalPort))
for internal, external := range internalToExternalPort {
portStrings = append(portStrings, portPairToString(internal, external))
}
sort.StringSlice(portStrings).Sort()
return "ports forwarded are " + strings.Join(portStrings[:len(portStrings)-1], ", ") +
" and " + portStrings[len(portStrings)-1]
}
}
func portPairToString(internal, external uint16) string {
if internal == external {
return strconv.FormatUint(uint64(external), 10)
}
return fmt.Sprintf("%d (internal port %d)", external, internal)
}
@@ -40,3 +40,42 @@ func Test_portsToString(t *testing.T) {
}) })
} }
} }
func Test_externalInternalPortsToString(t *testing.T) {
t.Parallel()
testCases := map[string]struct {
internalToExternalPort map[uint16]uint16
s string
}{
"no_port": {
s: "no port forwarded",
},
"one_port": {
internalToExternalPort: map[uint16]uint16{123: 123},
s: "port forwarded is 123",
},
"two_ports": {
internalToExternalPort: map[uint16]uint16{123: 123, 456: 456},
s: "ports forwarded are 123 and 456",
},
"two_ports_different_internal_external": {
internalToExternalPort: map[uint16]uint16{123: 124, 456: 457},
s: "ports forwarded are 124 (internal port 123) and 457 (internal port 456)",
},
"three_ports": {
internalToExternalPort: map[uint16]uint16{123: 123, 456: 456, 789: 789},
s: "ports forwarded are 123, 456 and 789",
},
}
for name, testCase := range testCases {
t.Run(name, func(t *testing.T) {
t.Parallel()
s := portPairsToString(testCase.internalToExternalPort)
assert.Equal(t, testCase.s, s)
})
}
}
+3 -5
View File
@@ -89,6 +89,9 @@ func (s *Service) Start(ctx context.Context) (runError <-chan error, err error)
} }
func (s *Service) onNewPorts(ctx context.Context, internalToExternalPorts map[uint16]uint16) (err error) { func (s *Service) onNewPorts(ctx context.Context, internalToExternalPorts map[uint16]uint16) (err error) {
s.logger.Info(portPairsToString(internalToExternalPorts))
externalPorts := slices.Collect(maps.Values(internalToExternalPorts))
autoRedirectionNeeded := false autoRedirectionNeeded := false
externalToInternalPorts := make(map[uint16]uint16, len(internalToExternalPorts)) externalToInternalPorts := make(map[uint16]uint16, len(internalToExternalPorts))
for internal, external := range internalToExternalPorts { for internal, external := range internalToExternalPorts {
@@ -97,12 +100,7 @@ func (s *Service) onNewPorts(ctx context.Context, internalToExternalPorts map[ui
autoRedirectionNeeded = true autoRedirectionNeeded = true
} }
} }
externalPorts := slices.Collect(maps.Keys(externalToInternalPorts))
slices.Sort(externalPorts) slices.Sort(externalPorts)
s.logger.Info(portsToString(externalPorts))
userRedirectionEnabled := !slices.Equal(s.settings.ListeningPorts, []uint16{0}) userRedirectionEnabled := !slices.Equal(s.settings.ListeningPorts, []uint16{0})
for i, port := range externalPorts { for i, port := range externalPorts {
internalPort := externalToInternalPorts[port] internalPort := externalToInternalPorts[port]
+1 -1
View File
@@ -73,7 +73,7 @@ func modifyConfig(lines []string, connection models.Connection,
modified = append(modified, "pull-filter ignore \"auth-token\"") // prevent auth failed loop modified = append(modified, "pull-filter ignore \"auth-token\"") // prevent auth failed loop
modified = append(modified, "auth-retry nointeract") modified = append(modified, "auth-retry nointeract")
modified = append(modified, "suppress-timestamps") modified = append(modified, "suppress-timestamps")
modified = append(modified, "hand-window 10") // default is 60 seconds which is too long modified = append(modified, "hand-window 20") // default is 60 seconds which is too long
if *settings.User != "" { if *settings.User != "" {
modified = append(modified, "auth-user-pass "+openvpn.AuthConf) modified = append(modified, "auth-user-pass "+openvpn.AuthConf)
} }
+1 -1
View File
@@ -62,7 +62,7 @@ func Test_modifyConfig(t *testing.T) {
"pull-filter ignore \"auth-token\"", "pull-filter ignore \"auth-token\"",
"auth-retry nointeract", "auth-retry nointeract",
"suppress-timestamps", "suppress-timestamps",
"hand-window 10", "hand-window 20",
"auth-user-pass /etc/openvpn/auth.conf", "auth-user-pass /etc/openvpn/auth.conf",
"verb 0", "verb 0",
"data-ciphers-fallback cipher", "data-ciphers-fallback cipher",
@@ -13,6 +13,7 @@ import (
"net/netip" "net/netip"
"net/url" "net/url"
"os" "os"
"regexp"
"strconv" "strconv"
"strings" "strings"
"time" "time"
@@ -79,13 +80,27 @@ func (p *Provider) PortForward(ctx context.Context,
} }
durationToExpiration = data.Expiration.Sub(p.timeNow()) durationToExpiration = data.Expiration.Sub(p.timeNow())
} }
logger.Info("Port forwarded data expires in " + format.FriendlyDuration(durationToExpiration))
// First time binding // First time binding
if err := bindPort(ctx, privateIPClient, p.apiIP, data); err != nil { for ctx.Err() == nil {
return nil, fmt.Errorf("binding port: %w", err) err = bindPort(ctx, privateIPClient, p.apiIP, data)
if err == nil {
break
} else if !errors.Is(err, errPortBusy) {
return nil, fmt.Errorf("binding port: %w", err)
}
logger.Warn("refreshing port forward data and trying again because " + err.Error())
client := objects.Client
data, err = refreshPIAPortForwardData(ctx, client, privateIPClient, p.apiIP,
p.portForwardPath, objects.Username, objects.Password)
if err != nil {
return nil, fmt.Errorf("refreshing port forward data: %w", err)
}
durationToExpiration = data.Expiration.Sub(p.timeNow())
} }
logger.Info("Port forwarded data expires in " + format.FriendlyDuration(durationToExpiration))
return map[uint16]uint16{data.Port: data.Port}, nil return map[uint16]uint16{data.Port: data.Port}, nil
} }
@@ -393,6 +408,13 @@ func fetchPortForwardData(ctx context.Context, client *http.Client, apiIP netip.
return port, data.Signature, expiration, err return port, data.Signature, expiration, err
} }
var errPortBusy = errors.New("port is busy")
var (
regexPortBusy = regexp.MustCompile(`^port \d+ is busy\. `)
regexNumber = regexp.MustCompile(`\d+`)
)
func bindPort(ctx context.Context, client *http.Client, apiIPAddress netip.Addr, data piaPortForwardData) (err error) { func bindPort(ctx context.Context, client *http.Client, apiIPAddress netip.Addr, data piaPortForwardData) (err error) {
// Define a timeout since the default client has a large timeout and we don't // Define a timeout since the default client has a large timeout and we don't
// want to wait too long. // want to wait too long.
@@ -431,7 +453,9 @@ func bindPort(ctx context.Context, client *http.Client, apiIPAddress netip.Addr,
} }
defer response.Body.Close() defer response.Body.Close()
if response.StatusCode != http.StatusOK { switch response.StatusCode {
case http.StatusOK, http.StatusConflict:
default:
return makeNOKStatusError(response, errSubstitutions) return makeNOKStatusError(response, errSubstitutions)
} }
@@ -444,11 +468,24 @@ func bindPort(ctx context.Context, client *http.Client, apiIPAddress netip.Addr,
return fmt.Errorf("decoding response: from %s: %w", bindPortURL.String(), err) return fmt.Errorf("decoding response: from %s: %w", bindPortURL.String(), err)
} }
if responseData.Status != "OK" { switch response.StatusCode {
return fmt.Errorf("bad response received with status %q and message %q", responseData.Status, responseData.Message) case http.StatusOK:
if responseData.Status != "OK" {
return fmt.Errorf("bad response received with status %q and message %q", responseData.Status, responseData.Message)
}
return nil
case http.StatusConflict:
portIsBusy := regexPortBusy.FindString(responseData.Message)
if portIsBusy == "" {
return fmt.Errorf("port busy response received with unexpected message %q not matching regex %q",
responseData.Message, regexPortBusy.String())
}
portStr := regexNumber.FindString(portIsBusy)
rest := strings.TrimPrefix(responseData.Message, portIsBusy)
return fmt.Errorf("%w: %s - %s", errPortBusy, portStr, rest)
default:
panic("unreachable code")
} }
return nil
} }
// replaceInErr is used to remove sensitive information from errors. // replaceInErr is used to remove sensitive information from errors.
@@ -470,6 +507,18 @@ func makeNOKStatusError(response *http.Response, substitutions map[string]string
url = replaceInString(url, substitutions) url = replaceInString(url, substitutions)
b, _ := io.ReadAll(response.Body) b, _ := io.ReadAll(response.Body)
var responseData struct {
Status string `json:"status"`
Message string `json:"message"`
}
if err := json.Unmarshal(b, &responseData); err == nil {
responseData.Message = replaceInString(responseData.Message, substitutions)
return fmt.Errorf("HTTP status code not OK: %s: %d %s: response received: status %q and message %q",
url, response.StatusCode, response.Status, responseData.Status, responseData.Message)
}
// Fallback on non JSON response body
shortenMessage := string(b) shortenMessage := string(b)
shortenMessage = strings.ReplaceAll(shortenMessage, "\n", "") shortenMessage = strings.ReplaceAll(shortenMessage, "\n", "")
shortenMessage = strings.ReplaceAll(shortenMessage, " ", " ") shortenMessage = strings.ReplaceAll(shortenMessage, " ", " ")
+43 -24
View File
@@ -16,22 +16,25 @@ import (
"strings" "strings"
srp "github.com/ProtonMail/go-srp" srp "github.com/ProtonMail/go-srp"
"github.com/qdm12/gluetun/internal/provider/common"
) )
// apiClient is a minimal Proton v4 API client which can handle all the // apiClient is a minimal Proton v4 API client which can handle all the
// oddities of Proton's authentication flow they want to keep hidden // oddities of Proton's authentication flow they want to keep hidden
// from the public. // from the public.
type apiClient struct { type apiClient struct {
apiURLBase string apiURLBase string
httpClient *http.Client httpClient *http.Client
appVersion string appVersion string
userAgent string vpnGtkAppVersion string
generator *rand.ChaCha8 userAgent string
generator *rand.ChaCha8
warner common.Warner
} }
// newAPIClient returns an [apiClient] with sane defaults matching Proton's // newAPIClient returns an [apiClient] with sane defaults matching Proton's
// insane expectations. // insane expectations.
func newAPIClient(ctx context.Context, httpClient *http.Client) (client *apiClient, err error) { func newAPIClient(ctx context.Context, httpClient *http.Client, warner common.Warner) (client *apiClient, err error) {
var seed [32]byte var seed [32]byte
_, _ = crand.Read(seed[:]) _, _ = crand.Read(seed[:])
generator := rand.NewChaCha8(seed) generator := rand.NewChaCha8(seed)
@@ -46,17 +49,23 @@ func newAPIClient(ctx context.Context, httpClient *http.Client) (client *apiClie
} }
userAgent := userAgents[generator.Uint64()%uint64(len(userAgents))] userAgent := userAgents[generator.Uint64()%uint64(len(userAgents))]
appVersion, err := getMostRecentStableTag(ctx, httpClient) appVersion, err := getMostRecentStableWebAccountTag(ctx, httpClient)
if err != nil { if err != nil {
return nil, fmt.Errorf("getting most recent version for proton app: %w", err) return nil, fmt.Errorf("getting most recent version for web-account: %w", err)
}
vpnGtkAppVersion, err := getMostRecentStableVPNGtkAppTag(ctx, httpClient)
if err != nil {
return nil, fmt.Errorf("getting most recent version for linux VPN GTK app: %w", err)
} }
return &apiClient{ return &apiClient{
apiURLBase: "https://account.proton.me/api", apiURLBase: "https://account.proton.me/api",
httpClient: httpClient, httpClient: httpClient,
appVersion: appVersion, appVersion: appVersion,
userAgent: userAgent, vpnGtkAppVersion: vpnGtkAppVersion,
generator: generator, userAgent: userAgent,
generator: generator,
warner: warner,
}, nil }, nil
} }
@@ -64,10 +73,10 @@ func newAPIClient(ctx context.Context, httpClient *http.Client) (client *apiClie
// to succeed without being blocked by their "security" measures. // to succeed without being blocked by their "security" measures.
// See for example [getMostRecentStableTag] on how the app version must // See for example [getMostRecentStableTag] on how the app version must
// be set to a recent version or they block your request. "SeCuRiTy"... // be set to a recent version or they block your request. "SeCuRiTy"...
func (c *apiClient) setHeaders(request *http.Request, cookie cookie) { func (c *apiClient) setHeaders(request *http.Request, cookie cookie, appVersion string) {
request.Header.Set("Cookie", cookie.String()) request.Header.Set("Cookie", cookie.String())
request.Header.Set("User-Agent", c.userAgent) request.Header.Set("User-Agent", c.userAgent)
request.Header.Set("x-pm-appversion", c.appVersion) request.Header.Set("x-pm-appversion", appVersion)
request.Header.Set("x-pm-locale", "en_US") request.Header.Set("x-pm-locale", "en_US")
request.Header.Set("x-pm-uid", cookie.uid) request.Header.Set("x-pm-uid", cookie.uid)
} }
@@ -98,7 +107,11 @@ func (c *apiClient) authenticate(ctx context.Context, email, password string,
} }
username, modulusPGPClearSigned, serverEphemeralBase64, saltBase64, username, modulusPGPClearSigned, serverEphemeralBase64, saltBase64,
srpSessionHex, version, err := c.authInfo(ctx, email, unauthCookie) srpSessionHex, version, err := c.authInfo(ctx, email, unauthCookie)
if err != nil { switch {
case errors.Is(err, errUsernameEmpty):
c.warner.Warn("Username is empty in auth info response, trying with email address instead")
username = email
case err != nil:
return cookie{}, fmt.Errorf("getting auth information: %w", err) return cookie{}, fmt.Errorf("getting auth information: %w", err)
} }
@@ -159,7 +172,7 @@ func (c *apiClient) getUnauthSession(ctx context.Context, sessionID string) (
unauthCookie := cookie{ unauthCookie := cookie{
sessionID: sessionID, sessionID: sessionID,
} }
c.setHeaders(request, unauthCookie) c.setHeaders(request, unauthCookie, c.appVersion)
response, err := c.httpClient.Do(request) response, err := c.httpClient.Do(request)
if err != nil { if err != nil {
@@ -244,7 +257,7 @@ func (c *apiClient) cookieToken(ctx context.Context, sessionID, tokenType, acces
uid: uid, uid: uid,
sessionID: sessionID, sessionID: sessionID,
} }
c.setHeaders(request, unauthCookie) c.setHeaders(request, unauthCookie, c.appVersion)
request.Header.Set("Authorization", tokenType+" "+accessToken) request.Header.Set("Authorization", tokenType+" "+accessToken)
response, err := c.httpClient.Do(request) response, err := c.httpClient.Do(request)
@@ -291,6 +304,8 @@ func (c *apiClient) cookieToken(ctx context.Context, sessionID, tokenType, acces
return "", errors.New("auth cookie not found") return "", errors.New("auth cookie not found")
} }
var errUsernameEmpty = errors.New("username is empty in response")
// authInfo fetches SRP parameters for the account. // authInfo fetches SRP parameters for the account.
func (c *apiClient) authInfo(ctx context.Context, email string, unauthCookie cookie) ( func (c *apiClient) authInfo(ctx context.Context, email string, unauthCookie cookie) (
username, modulusPGPClearSigned, serverEphemeralBase64, saltBase64, srpSessionHex string, username, modulusPGPClearSigned, serverEphemeralBase64, saltBase64, srpSessionHex string,
@@ -315,7 +330,7 @@ func (c *apiClient) authInfo(ctx context.Context, email string, unauthCookie coo
if err != nil { if err != nil {
return "", "", "", "", "", 0, fmt.Errorf("creating request: %w", err) return "", "", "", "", "", 0, fmt.Errorf("creating request: %w", err)
} }
c.setHeaders(request, unauthCookie) c.setHeaders(request, unauthCookie, c.appVersion)
request.Header.Set("Content-Type", "application/json") request.Header.Set("Content-Type", "application/json")
response, err := c.httpClient.Do(request) response, err := c.httpClient.Do(request)
@@ -358,15 +373,17 @@ func (c *apiClient) authInfo(ctx context.Context, email string, unauthCookie coo
return "", "", "", "", "", 0, errors.New("salt is empty in response") return "", "", "", "", "", 0, errors.New("salt is empty in response")
case info.SRPSession == "": case info.SRPSession == "":
return "", "", "", "", "", 0, errors.New("SRP session is empty in response") return "", "", "", "", "", 0, errors.New("SRP session is empty in response")
case info.Username == "":
return "", "", "", "", "", 0, errors.New("username is empty in response")
case info.Version == nil: case info.Version == nil:
return "", "", "", "", "", 0, errors.New("version is missing in response") return "", "", "", "", "", 0, errors.New("version is missing in response")
case info.Username == "":
// Return a sentinel error the caller can handle to try with the email address instead of the username.
// Some accounts seem to have no username.
err = fmt.Errorf("%w", errUsernameEmpty)
} }
version = int(*info.Version) //nolint:gosec version = int(*info.Version) //nolint:gosec
return info.Username, info.Modulus, info.ServerEphemeral, info.Salt, return info.Username, info.Modulus, info.ServerEphemeral, info.Salt,
info.SRPSession, version, nil info.SRPSession, version, err
} }
type cookie struct { type cookie struct {
@@ -422,7 +439,7 @@ func (c *apiClient) auth(ctx context.Context, unauthCookie cookie,
if err != nil { if err != nil {
return cookie{}, fmt.Errorf("creating request: %w", err) return cookie{}, fmt.Errorf("creating request: %w", err)
} }
c.setHeaders(request, unauthCookie) c.setHeaders(request, unauthCookie, c.appVersion)
request.Header.Set("Content-Type", "application/json") request.Header.Set("Content-Type", "application/json")
response, err := c.httpClient.Do(request) response, err := c.httpClient.Do(request)
@@ -573,7 +590,9 @@ func (c *apiClient) fetchServers(ctx context.Context, cookie cookie) (
if err != nil { if err != nil {
return data, err return data, err
} }
c.setHeaders(request, cookie) // Note we use the vpnGtkAppVersion field given it produces an output of more servers
c.setHeaders(request, cookie, c.vpnGtkAppVersion)
request.Header.Set("x-pm-appversion", "linux-vpn@4.15.2")
response, err := c.httpClient.Do(request) response, err := c.httpClient.Do(request)
if err != nil { if err != nil {
@@ -20,7 +20,7 @@ func (u *Updater) FetchServers(ctx context.Context, minServers int) (
return nil, fmt.Errorf("%w: password is empty", common.ErrCredentialsMissing) return nil, fmt.Errorf("%w: password is empty", common.ErrCredentialsMissing)
} }
apiClient, err := newAPIClient(ctx, u.client) apiClient, err := newAPIClient(ctx, u.client, u.warner)
if err != nil { if err != nil {
return nil, fmt.Errorf("creating API client: %w", err) return nil, fmt.Errorf("creating API client: %w", err)
} }
+47 -2
View File
@@ -7,15 +7,18 @@ import (
"io" "io"
"net/http" "net/http"
"regexp" "regexp"
"sort"
"strings" "strings"
"time" "time"
"golang.org/x/mod/semver"
) )
// getMostRecentStableTag finds the most recent proton-account stable tag version, // getMostRecentStableWebAccountTag finds the most recent proton-account stable tag version,
// in order to use it in the x-pm-appversion http request header. Because if we do // in order to use it in the x-pm-appversion http request header. Because if we do
// fall behind on versioning, Proton doesn't like it because they like to create // fall behind on versioning, Proton doesn't like it because they like to create
// complications where there is no need for it. Hence this function. // complications where there is no need for it. Hence this function.
func getMostRecentStableTag(ctx context.Context, client *http.Client) (version string, err error) { func getMostRecentStableWebAccountTag(ctx context.Context, client *http.Client) (version string, err error) {
page := 1 page := 1
regexVersion := regexp.MustCompile(`^proton-account@(\d+\.\d+\.\d+\.\d+)$`) regexVersion := regexp.MustCompile(`^proton-account@(\d+\.\d+\.\d+\.\d+)$`)
for ctx.Err() == nil { for ctx.Err() == nil {
@@ -69,3 +72,45 @@ func getMostRecentStableTag(ctx context.Context, client *http.Client) (version s
return "", fmt.Errorf("%w (queried %d pages)", context.Canceled, page) return "", fmt.Errorf("%w (queried %d pages)", context.Canceled, page)
} }
// getMostRecentStableVPNGtkAppTag finds the latest proton-vpn-gtk-app semver tag,
// in order to use it in the x-pm-appversion http request header ONLY to fetch servers
// data. Because if we do fall behind on versioning, Proton doesn't like it because they like
// to create complications where there is no need for it. Hence this function.
func getMostRecentStableVPNGtkAppTag(ctx context.Context, client *http.Client) (version string, err error) {
const url = "https://api.github.com/repos/ProtonVPN/proton-vpn-gtk-app/tags?per_page=30"
request, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return "", fmt.Errorf("creating request: %w", err)
}
request.Header.Set("Accept", "application/vnd.github.v3+json")
response, err := client.Do(request)
if err != nil {
return "", err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return "", fmt.Errorf("HTTP status code not OK: %s", response.Status)
}
decoder := json.NewDecoder(response.Body)
var data []struct {
Name string `json:"name"`
}
err = decoder.Decode(&data)
if err != nil {
return "", fmt.Errorf("decoding JSON response: %w", err)
}
// Sort tags by semver. Invalid tags are placed at the end and we ignore them.
// Yes, proton does push invalid semver tag names sometimes. Good job yet again.
sort.Slice(data, func(i, j int) bool {
return semver.Compare(data[i].Name, data[j].Name) > 0
})
version = "linux-vpn@" + data[0].Name[1:] // remove leading v
return version, nil
}
+7
View File
@@ -2,6 +2,7 @@ package utils
import ( import (
"fmt" "fmt"
"math/rand/v2"
"slices" "slices"
"github.com/qdm12/gluetun/internal/configuration/settings" "github.com/qdm12/gluetun/internal/configuration/settings"
@@ -44,6 +45,12 @@ func GetConnection(provider string,
return connection, fmt.Errorf("filtering servers: %w", err) return connection, fmt.Errorf("filtering servers: %w", err)
} }
// Randomize order of the servers struct so the first connection to be picked
// won't always be the same one.
rand.Shuffle(len(servers), func(i, j int) {
servers[i], servers[j] = servers[j], servers[i]
})
protocol := getProtocol(selection) protocol := getProtocol(selection)
port := getPort(selection, defaults.OpenVPNTCPPort, port := getPort(selection, defaults.OpenVPNTCPPort,
defaults.OpenVPNUDPPort, defaults.WireguardPort) defaults.OpenVPNUDPPort, defaults.WireguardPort)
+21 -19
View File
@@ -27,11 +27,12 @@ func Test_GetConnection(t *testing.T) {
defaults ConnectionDefaults defaults ConnectionDefaults
ipv6Supported bool ipv6Supported bool
randSource rand.Source randSource rand.Source
connection models.Connection connections []models.Connection
errMessage string errMessage string
}{ }{
"storage filter error": { "storage filter error": {
filterError: errors.New("test error"), filterError: errors.New("test error"),
connections: []models.Connection{{}},
errMessage: "filtering servers: test error", errMessage: "filtering servers: test error",
}, },
"server without IPs": { "server without IPs": {
@@ -46,7 +47,8 @@ func Test_GetConnection(t *testing.T) {
OpenVPNUDPPort: 1, OpenVPNUDPPort: 1,
WireguardPort: 1, WireguardPort: 1,
}, },
errMessage: "no connection to pick from", connections: []models.Connection{{}},
errMessage: "no connection to pick from",
}, },
"OpenVPN server with hostname": { "OpenVPN server with hostname": {
filteredServers: []models.Server{ filteredServers: []models.Server{
@@ -61,13 +63,13 @@ func Test_GetConnection(t *testing.T) {
WithDefaults(providers.Mullvad), WithDefaults(providers.Mullvad),
defaults: NewConnectionDefaults(443, 1194, 58820), defaults: NewConnectionDefaults(443, 1194, 58820),
randSource: rand.NewSource(0), randSource: rand.NewSource(0),
connection: models.Connection{ connections: []models.Connection{{
Type: vpn.OpenVPN, Type: vpn.OpenVPN,
IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}),
Protocol: constants.UDP, Protocol: constants.UDP,
Port: 1194, Port: 1194,
Hostname: "name", Hostname: "name",
}, }},
}, },
"OpenVPN server with x509": { "OpenVPN server with x509": {
filteredServers: []models.Server{ filteredServers: []models.Server{
@@ -83,13 +85,13 @@ func Test_GetConnection(t *testing.T) {
WithDefaults(providers.Mullvad), WithDefaults(providers.Mullvad),
defaults: NewConnectionDefaults(443, 1194, 58820), defaults: NewConnectionDefaults(443, 1194, 58820),
randSource: rand.NewSource(0), randSource: rand.NewSource(0),
connection: models.Connection{ connections: []models.Connection{{
Type: vpn.OpenVPN, Type: vpn.OpenVPN,
IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}),
Protocol: constants.UDP, Protocol: constants.UDP,
Port: 1194, Port: 1194,
Hostname: "x509", Hostname: "x509",
}, }},
}, },
"server with IPv4 and IPv6": { "server with IPv4 and IPv6": {
filteredServers: []models.Server{ filteredServers: []models.Server{
@@ -111,12 +113,12 @@ func Test_GetConnection(t *testing.T) {
WithDefaults(providers.Mullvad), WithDefaults(providers.Mullvad),
defaults: NewConnectionDefaults(443, 1194, 58820), defaults: NewConnectionDefaults(443, 1194, 58820),
randSource: rand.NewSource(0), randSource: rand.NewSource(0),
connection: models.Connection{ connections: []models.Connection{{
Type: vpn.OpenVPN, Type: vpn.OpenVPN,
IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}),
Protocol: constants.UDP, Protocol: constants.UDP,
Port: 1194, Port: 1194,
}, }},
}, },
"server with IPv4 and IPv6 and ipv6 supported": { "server with IPv4 and IPv6 and ipv6 supported": {
filteredServers: []models.Server{ filteredServers: []models.Server{
@@ -134,12 +136,12 @@ func Test_GetConnection(t *testing.T) {
defaults: NewConnectionDefaults(443, 1194, 58820), defaults: NewConnectionDefaults(443, 1194, 58820),
ipv6Supported: true, ipv6Supported: true,
randSource: rand.NewSource(0), randSource: rand.NewSource(0),
connection: models.Connection{ connections: []models.Connection{{
Type: vpn.OpenVPN, Type: vpn.OpenVPN,
IP: netip.IPv6Unspecified(), IP: netip.IPv6Unspecified(),
Protocol: constants.UDP, Protocol: constants.UDP,
Port: 1194, Port: 1194,
}, }},
}, },
"mixed servers": { "mixed servers": {
filteredServers: []models.Server{ filteredServers: []models.Server{
@@ -149,12 +151,6 @@ func Test_GetConnection(t *testing.T) {
IPs: []netip.Addr{netip.AddrFrom4([4]byte{1, 1, 1, 1})}, IPs: []netip.Addr{netip.AddrFrom4([4]byte{1, 1, 1, 1})},
OvpnX509: "ovpnx509", OvpnX509: "ovpnx509",
}, },
{
VPN: vpn.Wireguard,
UDP: true,
IPs: []netip.Addr{netip.AddrFrom4([4]byte{2, 2, 2, 2})},
OvpnX509: "ovpnx509",
},
{ {
VPN: vpn.OpenVPN, VPN: vpn.OpenVPN,
UDP: true, UDP: true,
@@ -169,13 +165,19 @@ func Test_GetConnection(t *testing.T) {
WithDefaults(providers.Mullvad), WithDefaults(providers.Mullvad),
defaults: NewConnectionDefaults(443, 1194, 58820), defaults: NewConnectionDefaults(443, 1194, 58820),
randSource: rand.NewSource(0), randSource: rand.NewSource(0),
connection: models.Connection{ connections: []models.Connection{{
Type: vpn.OpenVPN, Type: vpn.OpenVPN,
IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}),
Protocol: constants.UDP, Protocol: constants.UDP,
Port: 1194, Port: 1194,
Hostname: "ovpnx509", Hostname: "ovpnx509",
}, }, {
Type: vpn.OpenVPN,
IP: netip.AddrFrom4([4]byte{3, 3, 3, 3}),
Protocol: constants.UDP,
Port: 1194,
Hostname: "hostname",
}},
}, },
} }
@@ -194,7 +196,7 @@ func Test_GetConnection(t *testing.T) {
testCase.serverSelection, testCase.defaults, testCase.ipv6Supported, testCase.serverSelection, testCase.defaults, testCase.ipv6Supported,
connPicker) connPicker)
assert.Equal(t, testCase.connection, connection) assert.Contains(t, testCase.connections, connection)
if testCase.errMessage != "" { if testCase.errMessage != "" {
assert.EqualError(t, err, testCase.errMessage) assert.EqualError(t, err, testCase.errMessage)
} else { } else {
+1 -1
View File
@@ -62,7 +62,7 @@ func OpenVPNConfig(provider OpenVPNProviderSettings,
lines.add("mute-replay-warnings") // these are often ignored by some VPN providers lines.add("mute-replay-warnings") // these are often ignored by some VPN providers
lines.add("auth-retry", "nointeract") // retry authenticating without interaction lines.add("auth-retry", "nointeract") // retry authenticating without interaction
lines.add("suppress-timestamps") // do not log timestamps, the Gluetun logger takes care of it lines.add("suppress-timestamps") // do not log timestamps, the Gluetun logger takes care of it
lines.add("hand-window", "10") // default is 60 seconds which is too long lines.add("hand-window", "20") // default is 60 seconds which is too long
lines.add("dev", settings.Interface) lines.add("dev", settings.Interface)
lines.add("verb", fmt.Sprint(*settings.Verbosity)) lines.add("verb", fmt.Sprint(*settings.Verbosity))
protocol := connection.Protocol protocol := connection.Protocol
+63 -27
View File
@@ -2,6 +2,7 @@ package storage
import ( import (
"encoding/json" "encoding/json"
"fmt"
"os" "os"
"path/filepath" "path/filepath"
"sort" "sort"
@@ -9,44 +10,79 @@ import (
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
) )
// FlushToFile flushes the merged servers data to the file // flushToFile flushes the merged servers data to files
// specified by path, as indented JSON. // using the manifest file path given. It is not thread-safe.
func (s *Storage) FlushToFile(path string) error { func (s *Storage) flushToFile(manifestPath string) error {
s.mergedMutex.RLock() const (
defer s.mergedMutex.RUnlock() filePermission = 0o644
dirPermission = 0o755
)
return s.flushToFile(path) serversDirectoryPath := filepath.Dir(manifestPath)
} if err := os.MkdirAll(serversDirectoryPath, dirPermission); err != nil {
return fmt.Errorf("creating directory: %w", err)
// flushToFile flushes the merged servers data to the file
// specified by path, as indented JSON. It is not thread-safe.
func (s *Storage) flushToFile(path string) error {
if path == "" {
return nil // no file to write to
}
const permission = 0o644
dirPath := filepath.Dir(path)
if err := os.MkdirAll(dirPath, permission); err != nil {
return err
} }
file, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, permission) for provider, providerServers := range s.mergedServers.ProviderToServers {
providerFilepath := providerServers.Filepath
if providerFilepath == "" {
providerFilepath = filepath.Join(serversDirectoryPath, provider+".json")
}
providerDirectoryPath := filepath.Dir(providerFilepath)
if err := os.MkdirAll(providerDirectoryPath, dirPermission); err != nil {
return fmt.Errorf("creating directory: %w", err)
}
}
metadata := map[string]any{"version": s.mergedServers.Version}
for provider, providerServers := range s.mergedServers.ProviderToServers {
sort.Sort(models.SortableServers(providerServers.Servers))
providerFilepath := providerServers.Filepath
if providerFilepath == "" {
providerFilepath = filepath.Join(serversDirectoryPath, provider+".json")
}
providerFile, err := os.OpenFile(providerFilepath,
os.O_CREATE|os.O_WRONLY|os.O_TRUNC, filePermission)
if err != nil {
return fmt.Errorf("opening servers data file for %s: %w", provider, err)
}
encodedProviderServers := providerServers
encodedProviderServers.Filepath = ""
encoder := json.NewEncoder(providerFile)
encoder.SetIndent("", " ")
err = encoder.Encode(encodedProviderServers)
if err != nil {
_ = providerFile.Close()
return fmt.Errorf("encoding servers data for %s: %w", provider, err)
}
err = providerFile.Close()
if err != nil {
return fmt.Errorf("closing servers data file for %s: %w", provider, err)
}
metadata[provider] = map[string]string{"filepath": providerFilepath}
}
serversFile, err := os.OpenFile(manifestPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, filePermission)
if err != nil { if err != nil {
return err return err
} }
encoder := json.NewEncoder(file) encoder := json.NewEncoder(serversFile)
encoder.SetIndent("", " ") encoder.SetIndent("", " ")
for _, obj := range s.mergedServers.ProviderToServers { err = encoder.Encode(metadata)
sort.Sort(models.SortableServers(obj.Servers))
}
err = encoder.Encode(&s.mergedServers)
if err != nil { if err != nil {
_ = file.Close() _ = serversFile.Close()
return err return err
} }
return file.Close() return serversFile.Close()
} }
+64
View File
@@ -0,0 +1,64 @@
package storage
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"github.com/qdm12/gluetun/internal/models"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func Test_flushToFile_filepathOnlyInManifest(t *testing.T) {
t.Parallel()
tempPath := t.TempDir()
providerFilepath := filepath.Join(tempPath, "provider.json")
manifestPath := filepath.Join(tempPath, "manifest.json")
storage := &Storage{
mergedServers: models.AllServers{
Version: 1,
ProviderToServers: map[string]models.Servers{
"provider": {
Version: 1,
Timestamp: 1,
Filepath: providerFilepath,
},
},
},
}
err := storage.flushToFile(manifestPath)
require.NoError(t, err)
providerFile, err := os.Open(providerFilepath)
require.NoError(t, err)
defer providerFile.Close()
providerContent := make(map[string]json.RawMessage)
err = json.NewDecoder(providerFile).Decode(&providerContent)
require.NoError(t, err)
_, hasFilepath := providerContent["filepath"]
assert.False(t, hasFilepath)
manifestFile, err := os.Open(manifestPath)
require.NoError(t, err)
defer manifestFile.Close()
manifestContent := make(map[string]json.RawMessage)
err = json.NewDecoder(manifestFile).Decode(&manifestContent)
require.NoError(t, err)
providerMetadataRaw, ok := manifestContent["provider"]
require.True(t, ok)
var providerMetadata struct {
Filepath string `json:"filepath"`
}
err = json.Unmarshal(providerMetadataRaw, &providerMetadata)
require.NoError(t, err)
assert.Equal(t, providerFilepath, providerMetadata.Filepath)
}
+34 -3
View File
@@ -3,19 +3,50 @@ package storage
import ( import (
"embed" "embed"
"encoding/json" "encoding/json"
"fmt"
"path"
serversmodule "github.com/qdm12/gluetun-servers/pkg/servers"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
) )
//go:embed servers.json //go:embed servers.json
var allServersEmbedFS embed.FS var allServersEmbedFS embed.FS
func parseHardcodedServers() (allServers models.AllServers, err error) { func parseHardcodedServers() (allServers models.AllServers) {
f, err := allServersEmbedFS.Open("servers.json") f, err := allServersEmbedFS.Open("servers.json")
if err != nil { if err != nil {
return allServers, err panic(err)
} }
defer f.Close() // no-op
decoder := json.NewDecoder(f) decoder := json.NewDecoder(f)
err = decoder.Decode(&allServers) err = decoder.Decode(&allServers)
return allServers, err if err != nil {
panic("decoding servers.json: " + err.Error())
}
for provider, metadata := range allServers.ProviderToServers {
filename := path.Base(metadata.Filepath)
providerFile, err := serversmodule.Files.Open(filename)
if err != nil {
panic(fmt.Sprintf("reading embedded provider file %s for %s: %s", filename, provider, err))
}
defer providerFile.Close() // no-op
var providerServers models.Servers
decoder := json.NewDecoder(providerFile)
err = decoder.Decode(&providerServers)
if err != nil {
panic(fmt.Sprintf("JSON decoding embedded provider file %s for %s: %s",
filename, provider, err))
} else if providerServers.Filepath != "" {
panic(fmt.Sprintf("embedded provider file %s for %s should not have filepath set",
filename, provider))
}
providerServers.Filepath = metadata.Filepath // inherit filepath from servers.json
allServers.ProviderToServers[provider] = providerServers
}
return allServers
} }
+40 -3
View File
@@ -1,9 +1,13 @@
package storage package storage
import ( import (
"encoding/json"
"path"
"testing" "testing"
"github.com/qdm12/gluetun-servers/pkg/servers"
"github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/constants/providers"
"github.com/qdm12/gluetun/internal/models"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
@@ -11,9 +15,10 @@ import (
func Test_parseHardcodedServers(t *testing.T) { func Test_parseHardcodedServers(t *testing.T) {
t.Parallel() t.Parallel()
servers, err := parseHardcodedServers() var servers models.AllServers
assert.NotPanics(t, func() {
require.NoError(t, err) servers = parseHardcodedServers()
})
// all providers minus custom // all providers minus custom
allProviders := providers.All() allProviders := providers.All()
@@ -24,3 +29,35 @@ func Test_parseHardcodedServers(t *testing.T) {
assert.NotEmptyf(t, servers, "for provider %s", provider) assert.NotEmptyf(t, servers, "for provider %s", provider)
} }
} }
func Test_parseHardcodedServers_filepathsAndEmbeddedProviderFiles(t *testing.T) {
t.Parallel()
hardcodedServers := parseHardcodedServers()
allProviders := providers.All()
for _, provider := range allProviders {
providerServers, ok := hardcodedServers.ProviderToServers[provider]
require.Truef(t, ok, "for provider %s", provider)
require.NotEmptyf(t, providerServers.Filepath,
"embedded servers filepath should be set for provider %s", provider)
filename := path.Base(providerServers.Filepath)
file, err := servers.Files.Open(filename)
require.NoErrorf(t, err, "opening embedded provider file for %s", provider)
var fileServers struct {
Version uint16 `json:"version"`
Timestamp int64 `json:"timestamp"`
Servers []json.RawMessage `json:"servers"`
}
err = json.NewDecoder(file).Decode(&fileServers)
require.NoErrorf(t, err, "decoding embedded provider file for %s", provider)
require.NoError(t, file.Close())
assert.NotZerof(t, fileServers.Version, "for provider %s", provider)
assert.NotZerof(t, fileServers.Timestamp, "for provider %s", provider)
assert.NotEmptyf(t, fileServers.Servers, "for provider %s", provider)
}
}
+14
View File
@@ -30,6 +30,20 @@ func (s *Storage) mergeServers(hardcoded, persisted models.AllServers) models.Al
func (s *Storage) mergeProviderServers(provider string, func (s *Storage) mergeProviderServers(provider string,
hardcoded, persisted models.Servers, hardcoded, persisted models.Servers,
) (merged models.Servers) { ) (merged models.Servers) {
if persisted.Preferred && persisted.Version != hardcoded.Version {
s.logger.Warn(fmt.Sprintf(
"persisted preferred %s servers are discarded because they have version %d and hardcoded servers have version %d",
provider, persisted.Version, hardcoded.Version))
}
// If persisted data is marked as preferred, use it regardless of timestamp
// (as long as versions match)
if persisted.Preferred && persisted.Version == hardcoded.Version && len(persisted.Servers) > 0 {
s.logger.Info(fmt.Sprintf(
"Using %s servers from file (marked as preferred)", provider))
return persisted
}
nowTimestamp := time.Now().Unix() nowTimestamp := time.Now().Unix()
if persisted.Timestamp > nowTimestamp { if persisted.Timestamp > nowTimestamp {
s.logger.Warn(fmt.Sprintf( s.logger.Warn(fmt.Sprintf(
+17
View File
@@ -45,6 +45,23 @@ func (mr *MockLoggerMockRecorder) Info(arg0 interface{}) *gomock.Call {
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Info", reflect.TypeOf((*MockLogger)(nil).Info), arg0) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Info", reflect.TypeOf((*MockLogger)(nil).Info), arg0)
} }
// Infof mocks base method.
func (m *MockLogger) Infof(arg0 string, arg1 ...interface{}) {
m.ctrl.T.Helper()
varargs := []interface{}{arg0}
for _, a := range arg1 {
varargs = append(varargs, a)
}
m.ctrl.Call(m, "Infof", varargs...)
}
// Infof indicates an expected call of Infof.
func (mr *MockLoggerMockRecorder) Infof(arg0 interface{}, arg1 ...interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
varargs := append([]interface{}{arg0}, arg1...)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Infof", reflect.TypeOf((*MockLogger)(nil).Infof), varargs...)
}
// Warn mocks base method. // Warn mocks base method.
func (m *MockLogger) Warn(arg0 string) { func (m *MockLogger) Warn(arg0 string) {
m.ctrl.T.Helper() m.ctrl.T.Helper()
+85 -22
View File
@@ -12,29 +12,36 @@ import (
"golang.org/x/text/language" "golang.org/x/text/language"
) )
// readFromFile reads the servers from server.json. // readFromFile reads the servers data starting from the given manifest file path.
// It only reads servers that have the same version as the hardcoded servers version // It only reads servers that have the same version as the hardcoded servers version
// to avoid JSON decoding errors. // to avoid JSON decoding errors.
func (s *Storage) readFromFile(filepath string, hardcodedVersions map[string]uint16) ( func (s *Storage) readFromFile(manifestPath string, hardcodedVersions map[string]uint16) (
servers models.AllServers, err error, servers models.AllServers, found bool, err error,
) { ) {
file, err := os.Open(filepath) file, err := os.Open(manifestPath)
if os.IsNotExist(err) { if os.IsNotExist(err) {
return servers, nil return servers, false, nil
} else if err != nil { } else if err != nil {
return servers, err return servers, false, err
} }
b, err := io.ReadAll(file) b, err := io.ReadAll(file)
if err != nil { if err != nil {
return servers, err return servers, true, err
} }
if err := file.Close(); err != nil { if err := file.Close(); err != nil {
return servers, err return servers, true, err
} }
return s.extractServersFromBytes(b, hardcodedVersions) if len(b) == 0 {
// To satisfy https://github.com/qdm12/gluetun/issues/3318
// not too sure why but eh I'm feeling generous adding 3 code lines for it.
return servers, false, nil
}
servers, err = s.extractServersFromBytes(b, hardcodedVersions)
return servers, true, err
} }
func (s *Storage) extractServersFromBytes(b []byte, hardcodedVersions map[string]uint16) ( func (s *Storage) extractServersFromBytes(b []byte, hardcodedVersions map[string]uint16) (
@@ -46,6 +53,12 @@ func (s *Storage) extractServersFromBytes(b []byte, hardcodedVersions map[string
} }
// Note schema version is at map key "version" as number // Note schema version is at map key "version" as number
if rawVersion, ok := rawMessages["version"]; ok {
err := json.Unmarshal(rawVersion, &servers.Version)
if err != nil {
return servers, fmt.Errorf("decoding servers schema version: %w", err)
}
}
allProviders := providers.All() allProviders := providers.All()
servers.ProviderToServers = make(map[string]models.Servers, len(allProviders)) servers.ProviderToServers = make(map[string]models.Servers, len(allProviders))
@@ -86,25 +99,20 @@ func (s *Storage) readServers(provider string, hardcodedVersion uint16,
) { ) {
provider = titleCaser.String(provider) provider = titleCaser.String(provider)
var versionObject struct { var metadata struct {
Version uint16 `json:"version"` Version uint16 `json:"version"`
Timestamp uint64 `json:"timestamp"`
Filepath string `json:"filepath"`
} }
err = json.Unmarshal(rawMessage, &versionObject) err = json.Unmarshal(rawMessage, &metadata)
if err != nil { if err != nil {
return servers, false, fmt.Errorf("decoding servers version for provider %s: %w", return servers, false, fmt.Errorf("decoding servers version for provider %s: %w",
provider, err) provider, err)
} }
persistedVersion := versionObject.Version if metadata.Filepath != "" {
return s.readServersFromFilepath(provider, metadata.Filepath, hardcodedVersion)
versionsMatch = hardcodedVersion == persistedVersion
if !versionsMatch {
s.logger.Info(fmt.Sprintf(
"%s servers from file discarded because they have "+
"version %d and hardcoded servers have version %d",
provider, persistedVersion, hardcodedVersion))
return servers, versionsMatch, nil
} }
err = json.Unmarshal(rawMessage, &servers) err = json.Unmarshal(rawMessage, &servers)
@@ -113,5 +121,60 @@ func (s *Storage) readServers(provider string, hardcodedVersion uint16,
provider, err) provider, err)
} }
return servers, versionsMatch, nil const sourcePath = ""
if !checkVersions(hardcodedVersion, servers.Version, provider, sourcePath,
servers.Preferred, s.logger) {
return models.Servers{}, false, nil
}
return servers, true, nil
}
func (s *Storage) readServersFromFilepath(provider, filepath string, hardcodedVersion uint16) (
referencedServers models.Servers, versionsMatch bool, err error,
) {
providerFile, err := os.Open(filepath)
if os.IsNotExist(err) {
return models.Servers{}, false, nil
} else if err != nil {
return models.Servers{}, false, fmt.Errorf("opening servers file %s for provider %s: %w",
filepath, provider, err)
}
defer providerFile.Close()
err = json.NewDecoder(providerFile).Decode(&referencedServers)
if err != nil {
return models.Servers{}, false, fmt.Errorf("decoding servers file %s for provider %s: %w",
filepath, provider, err)
}
if !checkVersions(hardcodedVersion, referencedServers.Version, provider, filepath,
referencedServers.Preferred, s.logger) {
return models.Servers{}, false, nil
}
referencedServers.Filepath = filepath
return referencedServers, true, nil
}
func checkVersions(builtinVersion, version uint16, provider, sourcePath string,
preferred bool, logger Logger,
) (match bool) {
if version == builtinVersion {
return true
}
name := provider
log := logger.Info
if preferred {
name += " preferred"
log = logger.Warn
}
name += " servers"
if sourcePath != "" {
name += " from file " + sourcePath
}
log(fmt.Sprintf(
"%s discarded because they have version %d and hardcoded servers have version %d",
name, version, builtinVersion))
return false
} }
+32 -14
View File
@@ -30,25 +30,30 @@ func Test_extractServersFromBytes(t *testing.T) {
testCases := map[string]struct { testCases := map[string]struct {
b []byte b []byte
hardcodedVersions map[string]uint16 hardcodedVersions map[string]uint16
logged []string makeLogger func(ctrl *gomock.Controller) *MockLogger
persisted models.AllServers persisted models.AllServers
errMessage string errMessage string
}{ }{
"bad JSON": { "bad JSON": {
b: []byte("garbage"), b: []byte("garbage"),
makeLogger: func(_ *gomock.Controller) *MockLogger { return nil },
errMessage: "decoding servers: invalid character 'g' looking for beginning of value", errMessage: "decoding servers: invalid character 'g' looking for beginning of value",
}, },
"bad provider JSON": { "bad provider JSON": {
b: []byte(`{"cyberghost": "garbage"}`), b: []byte(`{"cyberghost": "garbage"}`),
hardcodedVersions: populateProviderToVersion(map[string]uint16{}), hardcodedVersions: populateProviderToVersion(map[string]uint16{}),
makeLogger: func(_ *gomock.Controller) *MockLogger { return nil },
errMessage: "decoding servers version for provider Cyberghost: " + errMessage: "decoding servers version for provider Cyberghost: " +
"json: cannot unmarshal string into Go value of type struct { Version uint16 \"json:\\\"version\\\"\" }", "json: cannot unmarshal string into Go value of type struct { Version uint16 \"json:\\\"version\\\"\"; " +
"Timestamp uint64 \"json:\\\"timestamp\\\"\"; " +
"Filepath string \"json:\\\"filepath\\\"\" }",
}, },
"bad servers array JSON": { "bad servers array JSON": {
b: []byte(`{"cyberghost": {"version": 1, "servers": "garbage"}}`), b: []byte(`{"cyberghost": {"version": 1, "servers": "garbage"}}`),
hardcodedVersions: populateProviderToVersion(map[string]uint16{ hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 1, providers.Cyberghost: 1,
}), }),
makeLogger: func(_ *gomock.Controller) *MockLogger { return nil },
errMessage: "decoding servers for provider Cyberghost: " + errMessage: "decoding servers for provider Cyberghost: " +
"json: cannot unmarshal string into Go struct field Servers.servers of type []models.Server", "json: cannot unmarshal string into Go struct field Servers.servers of type []models.Server",
}, },
@@ -57,6 +62,7 @@ func Test_extractServersFromBytes(t *testing.T) {
hardcodedVersions: populateProviderToVersion(map[string]uint16{ hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 1, providers.Cyberghost: 1,
}), }),
makeLogger: func(_ *gomock.Controller) *MockLogger { return nil },
persisted: models.AllServers{ persisted: models.AllServers{
ProviderToServers: map[string]models.Servers{}, ProviderToServers: map[string]models.Servers{},
}, },
@@ -68,6 +74,7 @@ func Test_extractServersFromBytes(t *testing.T) {
hardcodedVersions: populateProviderToVersion(map[string]uint16{ hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 1, providers.Cyberghost: 1,
}), }),
makeLogger: func(_ *gomock.Controller) *MockLogger { return nil },
persisted: models.AllServers{ persisted: models.AllServers{
ProviderToServers: map[string]models.Servers{ ProviderToServers: map[string]models.Servers{
providers.Cyberghost: {Version: 1}, providers.Cyberghost: {Version: 1},
@@ -81,8 +88,28 @@ func Test_extractServersFromBytes(t *testing.T) {
hardcodedVersions: populateProviderToVersion(map[string]uint16{ hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 2, providers.Cyberghost: 2,
}), }),
logged: []string{ makeLogger: func(ctrl *gomock.Controller) *MockLogger {
"Cyberghost servers from file discarded because they have version 1 and hardcoded servers have version 2", logger := NewMockLogger(ctrl)
logger.EXPECT().Info("Cyberghost servers discarded because " +
"they have version 1 and hardcoded servers have version 2")
return logger
},
persisted: models.AllServers{
ProviderToServers: map[string]models.Servers{},
},
},
"preferred_different_versions": {
b: []byte(`{
"cyberghost": {"version": 1, "timestamp": 1, "preferred": true}
}`),
hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 2,
}),
makeLogger: func(ctrl *gomock.Controller) *MockLogger {
logger := NewMockLogger(ctrl)
logger.EXPECT().Warn("Cyberghost preferred servers discarded because " +
"they have version 1 and hardcoded servers have version 2")
return logger
}, },
persisted: models.AllServers{ persisted: models.AllServers{
ProviderToServers: map[string]models.Servers{}, ProviderToServers: map[string]models.Servers{},
@@ -95,16 +122,7 @@ func Test_extractServersFromBytes(t *testing.T) {
t.Parallel() t.Parallel()
ctrl := gomock.NewController(t) ctrl := gomock.NewController(t)
logger := NewMockLogger(ctrl) logger := testCase.makeLogger(ctrl)
var previousLogCall *gomock.Call
for _, logged := range testCase.logged {
call := logger.EXPECT().Info(logged)
if previousLogCall != nil {
call.After(previousLogCall)
}
previousLogCall = call
}
s := &Storage{ s := &Storage{
logger: logger, logger: logger,
} }
+19 -3
View File
@@ -2,6 +2,8 @@ package storage
import ( import (
"fmt" "fmt"
"os"
"path/filepath"
"time" "time"
"github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/constants/providers"
@@ -10,12 +12,12 @@ import (
// SetServers sets the given servers for the given provider // SetServers sets the given servers for the given provider
// in the storage in-memory map and saves all the servers // in the storage in-memory map and saves all the servers
// to file. // to files.
// Note the servers given are not copied so the caller must // Note the servers given are not copied so the caller must
// NOT MUTATE them after calling this method. // NOT MUTATE them after calling this method.
func (s *Storage) SetServers(provider string, servers []models.Server) (err error) { func (s *Storage) SetServers(provider string, servers []models.Server) (err error) {
if provider == providers.Custom { if provider == providers.Custom {
return return nil
} }
s.mergedMutex.Lock() s.mergedMutex.Lock()
@@ -26,10 +28,24 @@ func (s *Storage) SetServers(provider string, servers []models.Server) (err erro
serversObject.Servers = servers serversObject.Servers = servers
s.mergedServers.ProviderToServers[provider] = serversObject s.mergedServers.ProviderToServers[provider] = serversObject
err = s.flushToFile(s.filepath) if !s.disk {
return nil
}
manifestPath := filepath.Join(s.directoryPath, manifestFilename)
err = s.flushToFile(manifestPath)
if err != nil { if err != nil {
return fmt.Errorf("saving servers to file: %w", err) return fmt.Errorf("saving servers to file: %w", err)
} }
if !s.hasLegacy() {
return nil
}
s.logger.Infof("removing legacy %s which is now migrated to %s", s.legacyFilepath, s.directoryPath)
err = os.Remove(s.legacyFilepath)
if err != nil {
s.logger.Warn("failed removing legacy servers file " + s.legacyFilepath + ": " + err.Error())
}
return nil return nil
} }
+23 -303854
View File
File diff suppressed because it is too large Load Diff
+32 -8
View File
@@ -1,6 +1,8 @@
package storage package storage
import ( import (
"os"
"path/filepath"
"sync" "sync"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
@@ -14,31 +16,38 @@ type Storage struct {
// SyncServers method. // SyncServers method.
hardcodedServers models.AllServers hardcodedServers models.AllServers
logger Logger logger Logger
filepath string disk bool
directoryPath string
legacyFilepath string
} }
const manifestFilename = "manifest.json"
type Logger interface { type Logger interface {
Info(s string) Info(s string)
Infof(format string, args ...any)
Warn(s string) Warn(s string)
} }
// New creates a new storage and reads the servers from the // New creates a new storage and reads the servers from the
// embedded servers file and the file on disk. // embedded servers files and the files on disk.
// Passing an empty filepath disables the reading and writing of // Passing an empty directoryPath disables the reading and writing of
// servers. // servers.
func New(logger Logger, filepath string) (storage *Storage, err error) { func New(logger Logger, disk bool, directoryPath, legacyFilepath string) (storage *Storage, err error) {
// A unit test prevents any error from being returned // A unit test prevents [parseHardcodedServers] from ever failing,
// and ensures all providers are part of the servers returned. // and ensures all providers are part of the servers returned.
hardcodedServers, _ := parseHardcodedServers() hardcodedServers := parseHardcodedServers()
storage = &Storage{ storage = &Storage{
hardcodedServers: hardcodedServers, hardcodedServers: hardcodedServers,
mergedServers: hardcodedServers, mergedServers: hardcodedServers,
logger: logger, logger: logger,
filepath: filepath, disk: disk,
directoryPath: directoryPath,
legacyFilepath: legacyFilepath,
} }
if filepath != "" { if disk {
if err := storage.syncServers(); err != nil { if err := storage.syncServers(); err != nil {
return nil, err return nil, err
} }
@@ -46,3 +55,18 @@ func New(logger Logger, filepath string) (storage *Storage, err error) {
return storage, nil return storage, nil
} }
// hasLegacy returns true if the legacy file `legacyFilepath` exists AND is
// different from the manifest file defined by `directoryPath`/[manifestFilename].
// This is used to determine if the legacy file should be read and removed after flushing servers data.
func (s *Storage) hasLegacy() bool {
if !s.disk {
return false
}
if filepath.Clean(filepath.Join(s.directoryPath, manifestFilename)) ==
filepath.Clean(s.legacyFilepath) {
return false
}
stat, err := os.Stat(s.legacyFilepath)
return err == nil && !stat.IsDir()
}
+32 -8
View File
@@ -2,6 +2,8 @@ package storage
import ( import (
"fmt" "fmt"
"os"
"path/filepath"
"reflect" "reflect"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
@@ -14,18 +16,31 @@ func countServers(allServers models.AllServers) (count int) {
return count return count
} }
// syncServers merges the hardcoded servers with the ones from the file. // syncServers merges the hardcoded servers with the ones from on disk files.
// It assumes s.directoryPath is set.
func (s *Storage) syncServers() (err error) { func (s *Storage) syncServers() (err error) {
hardcodedVersions := make(map[string]uint16, len(s.hardcodedServers.ProviderToServers)) hardcodedVersions := make(map[string]uint16, len(s.hardcodedServers.ProviderToServers))
for provider, servers := range s.hardcodedServers.ProviderToServers { for provider, servers := range s.hardcodedServers.ProviderToServers {
hardcodedVersions[provider] = servers.Version hardcodedVersions[provider] = servers.Version
} }
serversOnFile, err := s.readFromFile(s.filepath, hardcodedVersions) sourceManifestPath := filepath.Join(s.directoryPath, manifestFilename)
destinationManifestPath := sourceManifestPath
serversOnFile, found, err := s.readFromFile(sourceManifestPath, hardcodedVersions)
if err != nil { if err != nil {
return fmt.Errorf("reading servers from file: %w", err) return fmt.Errorf("reading servers from file: %w", err)
} }
hasLegacy := s.hasLegacy()
if !found && hasLegacy {
sourceManifestPath = s.legacyFilepath
s.logger.Infof("reading legacy servers file %s and migrating it to directory %s", sourceManifestPath, s.directoryPath)
serversOnFile, _, err = s.readFromFile(sourceManifestPath, hardcodedVersions)
if err != nil {
return fmt.Errorf("reading servers from file: %w", err)
}
}
hardcodedCount := countServers(s.hardcodedServers) hardcodedCount := countServers(s.hardcodedServers)
countOnFile := countServers(serversOnFile) countOnFile := countServers(serversOnFile)
@@ -34,13 +49,13 @@ func (s *Storage) syncServers() (err error) {
if countOnFile == 0 { if countOnFile == 0 {
s.logger.Info(fmt.Sprintf( s.logger.Info(fmt.Sprintf(
"creating %s with %d hardcoded servers", "writing servers data files to %s with %d hardcoded servers",
s.filepath, hardcodedCount)) s.directoryPath, hardcodedCount))
s.mergedServers = s.hardcodedServers s.mergedServers = s.hardcodedServers
} else { } else {
s.logger.Info(fmt.Sprintf( s.logger.Info(fmt.Sprintf(
"merging by most recent %d hardcoded servers and %d servers read from %s", "merging by most recent %d hardcoded servers and %d servers read from manifest file %s",
hardcodedCount, countOnFile, s.filepath)) hardcodedCount, countOnFile, sourceManifestPath))
s.mergedServers = s.mergeServers(s.hardcodedServers, serversOnFile) s.mergedServers = s.mergeServers(s.hardcodedServers, serversOnFile)
} }
@@ -50,9 +65,18 @@ func (s *Storage) syncServers() (err error) {
return nil return nil
} }
err = s.flushToFile(s.filepath) err = s.flushToFile(destinationManifestPath)
if err != nil { if err != nil {
s.logger.Warn("failed writing servers to file: " + err.Error()) s.logger.Warn("failed writing servers to destination manifest: " + err.Error())
return nil
}
migratedFromLegacy := hasLegacy && sourceManifestPath == s.legacyFilepath
if migratedFromLegacy {
err = os.Remove(sourceManifestPath)
if err != nil && !os.IsNotExist(err) {
s.logger.Warn("failed removing legacy servers file " + sourceManifestPath + ": " + err.Error())
}
} }
return nil return nil
} }
-41
View File
@@ -1,41 +0,0 @@
//go:build linux || darwin
package tun
import (
"errors"
"fmt"
"os"
"syscall"
)
// Check checks the tunnel device specified by path is present and accessible.
func (t *Tun) Check(path string) error {
f, err := os.OpenFile(path, os.O_RDWR, 0)
if err != nil {
return fmt.Errorf("TUN device is not available: %w", err)
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return fmt.Errorf("getting stat information for TUN file: %w", err)
}
sys, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return errors.New("cannot get syscall stat info of TUN file")
}
const expectedRdev = 2760 // corresponds to major 10 and minor 200
if sys.Rdev != expectedRdev {
return fmt.Errorf("TUN file has an unexpected rdev: %d instead of expected %d",
sys.Rdev, expectedRdev)
}
if err := f.Close(); err != nil {
return fmt.Errorf("closing TUN device: %w", err)
}
return nil
}
-7
View File
@@ -1,7 +0,0 @@
//go:build !linux && !darwin
package tun
func (t *Tun) Check(path string) error {
panic("not implemented")
}
-50
View File
@@ -1,50 +0,0 @@
//go:build linux || darwin
package tun
import (
"fmt"
"math"
"os"
"path/filepath"
"golang.org/x/sys/unix"
)
// Create creates a TUN device at the path specified.
func (t *Tun) Create(path string) (err error) {
parentDir := filepath.Dir(path)
err = os.MkdirAll(parentDir, 0o751) //nolint:mnd
if err != nil {
return err
}
const (
major = 10
minor = 200
)
dev := unix.Mkdev(major, minor)
if dev > math.MaxInt {
panic("dev is too high")
}
err = unix.Mknod(path, unix.S_IFCHR, int(dev))
if err != nil {
return fmt.Errorf("creating TUN device file node: %w", err)
}
fd, err := unix.Open(path, 0, 0)
if err != nil {
if err.Error() == "operation not permitted" {
err = fmt.Errorf("%w (did you specify --device /dev/net/tun to your container command?)", err)
}
return fmt.Errorf("unix opening TUN device file: %w", err)
}
const nonBlocking = true
err = unix.SetNonblock(fd, nonBlocking)
if err != nil {
return fmt.Errorf("setting non block to TUN device file descriptor: %w", err)
}
return nil
}
-8
View File
@@ -1,8 +0,0 @@
//go:build !linux && !darwin
package tun
// Create creates a TUN device at the path specified.
func (t *Tun) Create(path string) error {
panic("not implemented")
}
+96 -3
View File
@@ -1,7 +1,100 @@
//go:build linux || darwin
package tun package tun
type Tun struct{} import (
"errors"
"fmt"
"math"
"os"
"path/filepath"
"syscall"
func New() *Tun { "golang.org/x/sys/unix"
return &Tun{} )
func Setup() error {
const tunDevice = "/dev/net/tun"
err := check(tunDevice)
switch {
case err == nil:
return nil
case errors.Is(err, os.ErrNotExist):
err = create(tunDevice)
if err != nil {
return fmt.Errorf("creating TUN device: %w", err)
}
return nil
default:
return fmt.Errorf("checking TUN device: %w (see the Wiki errors/tun page)", err)
}
}
// check checks the tunnel device specified by path is present and accessible.
func check(path string) error {
f, err := os.OpenFile(path, os.O_RDWR, 0)
if err != nil {
return fmt.Errorf("TUN device is not available: %w", err)
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return fmt.Errorf("getting stat information for TUN file: %w", err)
}
sys, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return errors.New("cannot get syscall stat info of TUN file")
}
const expectedRdev = 2760 // corresponds to major 10 and minor 200
if sys.Rdev != expectedRdev {
return fmt.Errorf("TUN file has an unexpected rdev: %d instead of expected %d",
sys.Rdev, expectedRdev)
}
if err := f.Close(); err != nil {
return fmt.Errorf("closing TUN device: %w", err)
}
return nil
}
// create creates a TUN device at the path specified.
func create(path string) (err error) {
parentDir := filepath.Dir(path)
err = os.MkdirAll(parentDir, 0o751) //nolint:mnd
if err != nil {
return err
}
const (
major = 10
minor = 200
)
dev := unix.Mkdev(major, minor)
if dev > math.MaxInt {
panic("dev is too high")
}
err = unix.Mknod(path, unix.S_IFCHR, int(dev))
if err != nil {
return fmt.Errorf("creating TUN device file node: %w", err)
}
fd, err := unix.Open(path, 0, 0)
if err != nil {
if err.Error() == "operation not permitted" {
err = fmt.Errorf("%w (did you specify --device /dev/net/tun to your container command?)", err)
}
return fmt.Errorf("unix opening TUN device file: %w", err)
}
const nonBlocking = true
err = unix.SetNonblock(fd, nonBlocking)
if err != nil {
return fmt.Errorf("setting non block to TUN device file descriptor: %w", err)
}
return nil
} }
+6 -8
View File
@@ -10,20 +10,18 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func Test_Tun(t *testing.T) { func Test_Setup(t *testing.T) {
t.Parallel() t.Parallel()
path := getTempPath(t) path := getTempPath(t)
tun := New()
defer func() { defer func() {
err := os.RemoveAll(path) err := os.RemoveAll(path)
require.NoError(t, err) require.NoError(t, err)
}() }()
// No file check fail // No file check fail
err := tun.Check(path) err := check(path)
require.Error(t, err) require.Error(t, err)
expectedMessage := "TUN device is not available: open " + path + ": no such file or directory" expectedMessage := "TUN device is not available: open " + path + ": no such file or directory"
require.Equal(t, expectedMessage, err.Error()) require.Equal(t, expectedMessage, err.Error())
@@ -35,13 +33,13 @@ func Test_Tun(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Simple file check fail // Simple file check fail
err = tun.Check(path) err = check(path)
require.Error(t, err) require.Error(t, err)
expectedMessage = "TUN file has an unexpected rdev: 0 instead of expected 2760" expectedMessage = "TUN file has an unexpected rdev: 0 instead of expected 2760"
require.Equal(t, expectedMessage, err.Error()) require.Equal(t, expectedMessage, err.Error())
// Create TUN device fail as file exists // Create TUN device fail as file exists
err = tun.Create(path) err = create(path)
require.Error(t, err) require.Error(t, err)
require.EqualError(t, err, "creating TUN device file node: file exists") require.EqualError(t, err, "creating TUN device file node: file exists")
@@ -50,7 +48,7 @@ func Test_Tun(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Create TUN device success // Create TUN device success
err = tun.Create(path) err = create(path)
if err != nil && strings.HasSuffix(err.Error(), "operation not permitted") { if err != nil && strings.HasSuffix(err.Error(), "operation not permitted") {
t.Skip("You do not have root privileges to create a TUN device, skipping test") t.Skip("You do not have root privileges to create a TUN device, skipping test")
return return
@@ -58,7 +56,7 @@ func Test_Tun(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Check TUN device success // Check TUN device success
err = tun.Check(path) err = check(path)
require.NoError(t, err) require.NoError(t, err)
} }
+12
View File
@@ -0,0 +1,12 @@
//go:build !linux && !darwin
package tun
import (
"fmt"
"runtime"
)
func Setup() error {
return fmt.Errorf("not implemented for %s", runtime.GOOS)
}
+1 -1
View File
@@ -50,7 +50,7 @@ func NewLoop(settings settings.Updater, providers updater.Providers,
status: constants.Stopped, status: constants.Stopped,
settings: settings, settings: settings,
}, },
updater: updater.New(client, storage, providers, logger), updater: updater.New(client, storage, providers, logger, *settings.PreferDirectDownload),
logger: logger, logger: logger,
start: make(chan struct{}), start: make(chan struct{}),
running: make(chan models.LoopStatus), running: make(chan models.LoopStatus),
+36 -5
View File
@@ -5,6 +5,8 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"net/url"
"path"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
"github.com/qdm12/gluetun/internal/provider/common" "github.com/qdm12/gluetun/internal/provider/common"
@@ -16,18 +18,37 @@ type Provider interface {
} }
func (u *Updater) updateProvider(ctx context.Context, provider Provider, func (u *Updater) updateProvider(ctx context.Context, provider Provider,
minRatio float64, manifest manifest, minRatio float64,
) (err error) { ) (err error) {
providerName := provider.Name() providerName := provider.Name()
existingServersCount := u.storage.GetServersCount(providerName) existingServersCount := u.storage.GetServersCount(providerName)
minServers := int(minRatio * float64(existingServersCount)) minServers := int(minRatio * float64(existingServersCount))
servers, err := provider.FetchServers(ctx, minServers)
if err != nil { var servers []models.Server
if errors.Is(err, common.ErrNotEnoughServers) { if manifest.providerToFilepath == nil {
servers, err = provider.FetchServers(ctx, minServers)
switch {
case errors.Is(err, common.ErrNotEnoughServers):
u.logger.Warn("note: if running the update manually, you can use the flag " + u.logger.Warn("note: if running the update manually, you can use the flag " +
"-minratio to allow the update to succeed with less servers found") "-minratio to allow the update to succeed with less servers found")
fallthrough
case err != nil:
return fmt.Errorf("getting %s servers: %w", providerName, err)
}
} else {
providerFilepath := manifest.providerToFilepath[providerName]
providerFileURL := buildProviderFileURL(providerName, providerFilepath)
var data models.Servers
err = u.fetchJSON(ctx, providerFileURL, &data)
if err != nil {
return fmt.Errorf("downloading provider file %s: %w", providerFileURL, err)
}
servers = data.Servers
if len(servers) < minServers {
return fmt.Errorf("provider %s has not enough servers from downloaded file: got %d and expected at least %d",
providerName, len(servers), minServers)
} }
return fmt.Errorf("getting %s servers: %w", providerName, err)
} }
for _, server := range servers { for _, server := range servers {
@@ -55,3 +76,13 @@ func (u *Updater) updateProvider(ctx context.Context, provider Provider,
} }
return nil return nil
} }
func buildProviderFileURL(providerName, filePath string) (providerFileURL string) {
filename := path.Base(filePath)
if filename == "." || filename == "/" || filename == "" {
filename = providerName + ".json"
}
const serversFilesBaseURL = "https://raw.githubusercontent.com/qdm12/gluetun-servers/main/pkg/servers/"
return serversFilesBaseURL + url.PathEscape(filename)
}
+81 -9
View File
@@ -2,10 +2,15 @@ package updater
import ( import (
"context" "context"
"encoding/json"
"errors" "errors"
"fmt"
"io"
"net/http" "net/http"
"strings"
"time" "time"
"github.com/qdm12/gluetun/internal/constants/providers"
"github.com/qdm12/gluetun/internal/provider/common" "github.com/qdm12/gluetun/internal/provider/common"
"github.com/qdm12/gluetun/internal/updater/unzip" "github.com/qdm12/gluetun/internal/updater/unzip"
"golang.org/x/text/cases" "golang.org/x/text/cases"
@@ -13,7 +18,8 @@ import (
) )
type Updater struct { type Updater struct {
providers Providers providers Providers
preferDirectDownload bool
// state // state
storage Storage storage Storage
@@ -26,22 +32,31 @@ type Updater struct {
} }
func New(httpClient *http.Client, storage Storage, func New(httpClient *http.Client, storage Storage,
providers Providers, logger Logger, providers Providers, logger Logger, preferDirectDownload bool,
) *Updater { ) *Updater {
unzipper := unzip.New(httpClient) unzipper := unzip.New(httpClient)
return &Updater{ return &Updater{
providers: providers, providers: providers,
storage: storage, storage: storage,
logger: logger, logger: logger,
timeNow: time.Now, timeNow: time.Now,
client: httpClient, client: httpClient,
unzipper: unzipper, unzipper: unzipper,
preferDirectDownload: preferDirectDownload,
} }
} }
func (u *Updater) UpdateServers(ctx context.Context, providers []string, func (u *Updater) UpdateServers(ctx context.Context, providers []string,
minRatio float64, minRatio float64,
) (err error) { ) (err error) {
var manifest manifest
if u.preferDirectDownload {
manifest, err = u.fetchManifest(ctx)
if err != nil {
return fmt.Errorf("fetching remote manifest: %w", err)
}
}
caser := cases.Title(language.English) caser := cases.Title(language.English)
for _, providerName := range providers { for _, providerName := range providers {
u.logger.Info("updating " + caser.String(providerName) + " servers...") u.logger.Info("updating " + caser.String(providerName) + " servers...")
@@ -49,7 +64,7 @@ func (u *Updater) UpdateServers(ctx context.Context, providers []string,
fetcher := u.providers.Get(providerName) fetcher := u.providers.Get(providerName)
// TODO support servers offering only TCP or only UDP // TODO support servers offering only TCP or only UDP
// for NordVPN and PureVPN // for NordVPN and PureVPN
err := u.updateProvider(ctx, fetcher, minRatio) err := u.updateProvider(ctx, fetcher, manifest, minRatio)
switch { switch {
case err == nil: case err == nil:
continue continue
@@ -70,3 +85,60 @@ func (u *Updater) UpdateServers(ctx context.Context, providers []string,
return nil return nil
} }
type manifest struct {
providerToFilepath map[string]string
}
func (u *Updater) fetchManifest(ctx context.Context) (m manifest, err error) {
const serversManifestURL = "https://raw.githubusercontent.com/qdm12/gluetun-servers/main/pkg/servers/manifest.json"
var raw map[string]json.RawMessage
err = u.fetchJSON(ctx, serversManifestURL, &raw)
if err != nil {
return m, err
}
providerNames := providers.All()
m.providerToFilepath = make(map[string]string, len(providerNames))
for _, name := range providerNames {
var metadata struct {
Filepath string `json:"filepath"`
}
err = json.Unmarshal(raw[name], &metadata)
if err != nil {
return m, fmt.Errorf("decoding manifest metadata for %s: %w", name, err)
} else if metadata.Filepath == "" {
return m, fmt.Errorf("manifest missing filepath for provider %s", name)
}
m.providerToFilepath[name] = metadata.Filepath
}
return m, nil
}
func (u *Updater) fetchJSON(ctx context.Context, rawURL string, dst any) (err error) {
request, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
if err != nil {
return fmt.Errorf("creating request: %w", err)
}
response, err := u.client.Do(request)
if err != nil {
return fmt.Errorf("doing request: %w", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
const limit = 10 * 1024 * 1024 // 10 MiB
body, _ := io.ReadAll(io.LimitReader(response.Body, limit))
return fmt.Errorf("HTTP status code %d for %s: %s",
response.StatusCode, rawURL, strings.TrimSpace(string(body)))
}
err = json.NewDecoder(response.Body).Decode(dst)
if err != nil {
return fmt.Errorf("decoding response body: %w", err)
}
return nil
}
+6
View File
@@ -9,6 +9,7 @@ import (
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
"github.com/qdm12/gluetun/internal/netlink" "github.com/qdm12/gluetun/internal/netlink"
"github.com/qdm12/gluetun/internal/provider" "github.com/qdm12/gluetun/internal/provider"
"github.com/qdm12/gluetun/internal/tun"
"github.com/qdm12/gluetun/internal/wireguard" "github.com/qdm12/gluetun/internal/wireguard"
"github.com/qdm12/gosettings" "github.com/qdm12/gosettings"
) )
@@ -19,6 +20,11 @@ func setupAmneziaWg(ctx context.Context, netlinker NetLinker,
settings settings.VPN, ipv6SupportLevel netlink.IPv6SupportLevel, logger wireguard.Logger) ( settings settings.VPN, ipv6SupportLevel netlink.IPv6SupportLevel, logger wireguard.Logger) (
amneziawger *amneziawg.Amneziawg, connection models.Connection, err error, amneziawger *amneziawg.Amneziawg, connection models.Connection, err error,
) { ) {
err = tun.Setup()
if err != nil {
return nil, models.Connection{}, fmt.Errorf("setting up tun device: %w", err)
}
ipv6Internet := ipv6SupportLevel == netlink.IPv6Internet ipv6Internet := ipv6SupportLevel == netlink.IPv6Internet
connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet) connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet)
if err != nil { if err != nil {
+1 -2
View File
@@ -45,8 +45,6 @@ type Provider interface {
GetConnection(selection settings.ServerSelection, ipv6Supported bool) (connection models.Connection, err error) GetConnection(selection settings.ServerSelection, ipv6Supported bool) (connection models.Connection, err error)
OpenVPNConfig(connection models.Connection, settings settings.OpenVPN, ipv6Supported bool) (lines []string) OpenVPNConfig(connection models.Connection, settings settings.OpenVPN, ipv6Supported bool) (lines []string)
Name() string Name() string
FetchServers(ctx context.Context, minServers int) (
servers []models.Server, err error)
} }
type PortForwarder interface { type PortForwarder interface {
@@ -61,6 +59,7 @@ type Storage interface {
} }
type NetLinker interface { type NetLinker interface {
AddrList(linkIndex uint32, family uint8) (addresses []netip.Prefix, err error)
AddrReplace(linkIndex uint32, addr netip.Prefix) error AddrReplace(linkIndex uint32, addr netip.Prefix) error
Router Router
Ruler Ruler
+48
View File
@@ -0,0 +1,48 @@
package vpn
import (
"github.com/qdm12/gluetun/internal/configuration/settings"
"github.com/qdm12/gluetun/internal/constants/vpn"
"github.com/qdm12/gluetun/internal/netlink"
)
func (l *Loop) isIPv6Used(settings settings.VPN) bool {
if !l.ipv6SupportLevel.IsSupported() {
return false
}
switch settings.Type {
case vpn.AmneziaWg:
for _, prefix := range settings.AmneziaWg.Wireguard.Addresses {
if prefix.Addr().Is6() {
return true
}
}
return false
case vpn.OpenVPN:
link, err := l.netLinker.LinkByName(settings.OpenVPN.Interface)
if err != nil {
l.logger.Warnf("assuming IPv6 is not supported, cannot get OpenVPN link by name: %v", err)
return false
}
ipv6Prefixes, err := l.netLinker.AddrList(link.Index, netlink.FamilyV6)
if err != nil {
l.logger.Warnf("assuming IPv6 is not supported, cannot list OpenVPN link addresses: %v", err)
return false
}
for _, prefix := range ipv6Prefixes {
if prefix.Addr().IsGlobalUnicast() && !prefix.Addr().IsPrivate() {
return true
}
}
return false
case vpn.Wireguard:
for _, prefix := range settings.Wireguard.Addresses {
if prefix.Addr().Is6() {
return true
}
}
return false
default:
panic("vpn type not implemented: " + settings.Type)
}
}
+6
View File
@@ -9,6 +9,7 @@ import (
"github.com/qdm12/gluetun/internal/netlink" "github.com/qdm12/gluetun/internal/netlink"
"github.com/qdm12/gluetun/internal/openvpn" "github.com/qdm12/gluetun/internal/openvpn"
"github.com/qdm12/gluetun/internal/provider" "github.com/qdm12/gluetun/internal/provider"
"github.com/qdm12/gluetun/internal/tun"
) )
// setupOpenVPN sets OpenVPN up using the configurators and settings given. // setupOpenVPN sets OpenVPN up using the configurators and settings given.
@@ -18,6 +19,11 @@ func setupOpenVPN(ctx context.Context, fw Firewall,
settings settings.VPN, ipv6SupportLevel netlink.IPv6SupportLevel, starter Cmder, settings settings.VPN, ipv6SupportLevel netlink.IPv6SupportLevel, starter Cmder,
logger openvpn.Logger) (runner *openvpn.Runner, connection models.Connection, err error, logger openvpn.Logger) (runner *openvpn.Runner, connection models.Connection, err error,
) { ) {
err = tun.Setup()
if err != nil {
return nil, models.Connection{}, fmt.Errorf("setting up tun device: %w", err)
}
ipv6Internet := ipv6SupportLevel == netlink.IPv6Internet ipv6Internet := ipv6SupportLevel == netlink.IPv6Internet
connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet) connection, err = providerConf.GetConnection(settings.Provider.ServerSelection, ipv6Internet)
if err != nil { if err != nil {
+1
View File
@@ -59,6 +59,7 @@ func (l *Loop) Run(ctx context.Context, done chan<- struct{}) {
enabled: settings.Type != vpn.Wireguard || *settings.Wireguard.MTU == 0, enabled: settings.Type != vpn.Wireguard || *settings.Wireguard.MTU == 0,
vpnType: settings.Type, vpnType: settings.Type,
network: connection.Protocol, network: connection.Protocol,
ipv6: l.isIPv6Used(settings),
icmpAddrs: settings.PMTUD.ICMPAddresses, icmpAddrs: settings.PMTUD.ICMPAddresses,
tcpAddrs: settings.PMTUD.TCPAddresses, tcpAddrs: settings.PMTUD.TCPAddresses,
}, },
+15 -8
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"fmt" "fmt"
"net/netip" "net/netip"
"slices"
"strings" "strings"
"time" "time"
@@ -40,6 +41,8 @@ type tunnelUpPMTUDData struct {
// network is used to find the network level header overhead. // network is used to find the network level header overhead.
// It can be [constants.UDP] or [constants.TCP]. // It can be [constants.UDP] or [constants.TCP].
network string network string
// ipv6 is true if the VPN connection supports IPv6.
ipv6 bool
// icmpAddrs is the list of addresses to use for ICMP path MTU discovery. // icmpAddrs is the list of addresses to use for ICMP path MTU discovery.
// Each address should handle ICMP packets for PMTUD to work. // Each address should handle ICMP packets for PMTUD to work.
icmpAddrs []netip.Addr icmpAddrs []netip.Addr
@@ -69,7 +72,7 @@ func (l *Loop) onTunnelUp(ctx, loopCtx context.Context, data tunnelUpData) {
if data.pmtud.enabled { if data.pmtud.enabled {
mtuLogger := l.logger.New(log.SetComponent("MTU discovery")) mtuLogger := l.logger.New(log.SetComponent("MTU discovery"))
err := updateToMaxMTU(ctx, data.vpnIntf, data.pmtud.vpnType, err := updateToMaxMTU(ctx, data.vpnIntf, data.pmtud.vpnType,
data.pmtud.network, data.pmtud.icmpAddrs, data.pmtud.tcpAddrs, data.pmtud.network, data.pmtud.ipv6, data.pmtud.icmpAddrs, data.pmtud.tcpAddrs,
l.netLinker, l.routing, l.fw, mtuLogger) l.netLinker, l.routing, l.fw, mtuLogger)
if err != nil { if err != nil {
mtuLogger.Error(err.Error()) mtuLogger.Error(err.Error())
@@ -173,16 +176,11 @@ func (l *Loop) restartVPN(ctx context.Context, healthErr error) {
} }
func updateToMaxMTU(ctx context.Context, vpnInterface string, func updateToMaxMTU(ctx context.Context, vpnInterface string,
vpnType, network string, icmpAddrs []netip.Addr, tcpAddrs []netip.AddrPort, vpnType, network string, ipv6 bool, icmpAddrs []netip.Addr, tcpAddrs []netip.AddrPort,
netlinker NetLinker, routing Routing, firewall tcp.Firewall, logger *log.Logger, netlinker NetLinker, routing Routing, firewall tcp.Firewall, logger *log.Logger,
) error { ) error {
logger.Info("finding maximum MTU, this can take up to 6 seconds") logger.Info("finding maximum MTU, this can take up to 6 seconds")
vpnGatewayIP, err := routing.VPNLocalGatewayIP(vpnInterface)
if err != nil {
return fmt.Errorf("getting VPN gateway IP address: %w", err)
}
vpnRoutes, err := routing.VPNRoutes(vpnInterface) vpnRoutes, err := routing.VPNRoutes(vpnInterface)
if err != nil { if err != nil {
return fmt.Errorf("getting VPN routes: %w", err) return fmt.Errorf("getting VPN routes: %w", err)
@@ -195,7 +193,7 @@ func updateToMaxMTU(ctx context.Context, vpnInterface string,
originalMTU := link.MTU originalMTU := link.MTU
vpnLinkMTU := pmtud.MaxTheoreticalVPNMTU(vpnType, network, vpnGatewayIP) vpnLinkMTU := pmtud.MaxTheoreticalVPNMTU(vpnType, network, ipv6)
// Setting the VPN link MTU to 1500 might interrupt the connection until // Setting the VPN link MTU to 1500 might interrupt the connection until
// the new MTU is set again, but this is necessary to find the highest valid MTU. // the new MTU is set again, but this is necessary to find the highest valid MTU.
@@ -206,6 +204,15 @@ func updateToMaxMTU(ctx context.Context, vpnInterface string,
return fmt.Errorf("setting VPN interface %s MTU to %d: %w", vpnInterface, vpnLinkMTU, err) return fmt.Errorf("setting VPN interface %s MTU to %d: %w", vpnInterface, vpnLinkMTU, err)
} }
if !ipv6 {
icmpAddrs = slices.DeleteFunc(icmpAddrs, func(addr netip.Addr) bool {
return addr.Is6()
})
tcpAddrs = slices.DeleteFunc(tcpAddrs, func(addr netip.AddrPort) bool {
return addr.Addr().Is6()
})
}
const pingTimeout = time.Second const pingTimeout = time.Second
vpnLinkMTU, err = pmtud.PathMTUDiscover(ctx, icmpAddrs, tcpAddrs, vpnLinkMTU, err = pmtud.PathMTUDiscover(ctx, icmpAddrs, tcpAddrs,
vpnLinkMTU, pingTimeout, firewall, logger) vpnLinkMTU, pingTimeout, firewall, logger)
+12
View File
@@ -8,6 +8,7 @@ import (
"github.com/qdm12/gluetun/internal/cleanup" "github.com/qdm12/gluetun/internal/cleanup"
"github.com/qdm12/gluetun/internal/netlink" "github.com/qdm12/gluetun/internal/netlink"
gtun "github.com/qdm12/gluetun/internal/tun"
"golang.zx2c4.com/wireguard/conn" "golang.zx2c4.com/wireguard/conn"
"golang.zx2c4.com/wireguard/device" "golang.zx2c4.com/wireguard/device"
"golang.zx2c4.com/wireguard/tun" "golang.zx2c4.com/wireguard/tun"
@@ -27,15 +28,18 @@ func (w *Wireguard) Run(ctx context.Context, waitError chan<- error, ready chan<
} }
setupFunction := setupUserSpace setupFunction := setupUserSpace
userspace := false
switch w.settings.Implementation { switch w.settings.Implementation {
case "auto": //nolint:goconst case "auto": //nolint:goconst
if !kernelSupported { if !kernelSupported {
w.logger.Info("Using userspace implementation since Kernel support does not exist") w.logger.Info("Using userspace implementation since Kernel support does not exist")
userspace = true
break break
} }
w.logger.Info("Using available kernelspace implementation") w.logger.Info("Using available kernelspace implementation")
setupFunction = setupKernelSpace setupFunction = setupKernelSpace
case "userspace": case "userspace":
userspace = true
case "kernelspace": case "kernelspace":
if !kernelSupported { if !kernelSupported {
waitError <- errors.New("kernel does not support Wireguard") waitError <- errors.New("kernel does not support Wireguard")
@@ -46,6 +50,14 @@ func (w *Wireguard) Run(ctx context.Context, waitError chan<- error, ready chan<
panic(fmt.Sprintf("unknown implementation %q", w.settings.Implementation)) panic(fmt.Sprintf("unknown implementation %q", w.settings.Implementation))
} }
if userspace {
err = gtun.Setup()
if err != nil {
waitError <- fmt.Errorf("setting up userspace tun device: %w", err)
return
}
}
setup := func(ctx context.Context, cleanups *cleanup.Cleanups) ( setup := func(ctx context.Context, cleanups *cleanup.Cleanups) (
linkIndex uint32, waitAndCleanup func() error, err error, linkIndex uint32, waitAndCleanup func() error, err error,
) { ) {