diff --git a/internal/storage/read.go b/internal/storage/read.go index cc37a998..fed995b3 100644 --- a/internal/storage/read.go +++ b/internal/storage/read.go @@ -12,29 +12,30 @@ import ( "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 // to avoid JSON decoding errors. -func (s *Storage) readFromFile(filepath string, hardcodedVersions map[string]uint16) ( - servers models.AllServers, err error, +func (s *Storage) readFromFile(manifestPath string, hardcodedVersions map[string]uint16) ( + servers models.AllServers, found bool, err error, ) { - file, err := os.Open(filepath) + file, err := os.Open(manifestPath) if os.IsNotExist(err) { - return servers, nil + return servers, false, nil } else if err != nil { - return servers, err + return servers, false, err } b, err := io.ReadAll(file) if err != nil { - return servers, err + return servers, true, err } if err := file.Close(); err != nil { - return servers, err + return servers, true, err } - return s.extractServersFromBytes(b, hardcodedVersions) + servers, err = s.extractServersFromBytes(b, hardcodedVersions) + return servers, true, err } func (s *Storage) extractServersFromBytes(b []byte, hardcodedVersions map[string]uint16) ( @@ -46,6 +47,12 @@ func (s *Storage) extractServersFromBytes(b []byte, hardcodedVersions map[string } // 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() servers.ProviderToServers = make(map[string]models.Servers, len(allProviders)) @@ -86,25 +93,20 @@ func (s *Storage) readServers(provider string, hardcodedVersion uint16, ) { provider = titleCaser.String(provider) - var versionObject struct { - Version uint16 `json:"version"` + var metadata struct { + Version uint16 `json:"version"` + Timestamp int64 `json:"timestamp"` + Filepath string `json:"filepath"` } - err = json.Unmarshal(rawMessage, &versionObject) + err = json.Unmarshal(rawMessage, &metadata) if err != nil { return servers, false, fmt.Errorf("decoding servers version for provider %s: %w", provider, err) } - persistedVersion := versionObject.Version - - 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 + if metadata.Filepath != "" { + return s.readServersFromFilepath(provider, metadata.Filepath, hardcodedVersion) } err = json.Unmarshal(rawMessage, &servers) @@ -113,5 +115,55 @@ func (s *Storage) readServers(provider string, hardcodedVersion uint16, provider, err) } - return servers, versionsMatch, nil + versionsMatch = servers.Version == hardcodedVersion + if !versionsMatch { + if servers.Preferred { + s.logger.Warn(fmt.Sprintf( + "%s preferred servers from file discarded because they have version %d and hardcoded servers have version %d", + provider, servers.Version, hardcodedVersion)) + } else { + s.logger.Info(fmt.Sprintf( + "%s servers from file discarded because they have version %d and hardcoded servers have version %d", + provider, servers.Version, hardcodedVersion)) + } + return 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) + } + + versionsMatch = referencedServers.Version == hardcodedVersion + if !versionsMatch { + if referencedServers.Preferred { + s.logger.Warn(fmt.Sprintf( + "%s preferred servers from file %s discarded because they have version %d and hardcoded servers have version %d", + provider, filepath, referencedServers.Version, hardcodedVersion)) + } else { + s.logger.Info(fmt.Sprintf( + "%s servers from file %s discarded because they have version %d and hardcoded servers have version %d", + provider, filepath, referencedServers.Version, hardcodedVersion)) + } + return models.Servers{}, false, nil + } + + referencedServers.Filepath = filepath + return referencedServers, true, nil } diff --git a/internal/storage/read_test.go b/internal/storage/read_test.go index 20bbe17a..8be4f2ad 100644 --- a/internal/storage/read_test.go +++ b/internal/storage/read_test.go @@ -7,6 +7,7 @@ import ( "github.com/golang/mock/gomock" "github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/models" + "github.com/qdm12/log" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -27,10 +28,15 @@ func populateProviderToVersion(providerToVersion map[string]uint16) map[string]u func Test_extractServersFromBytes(t *testing.T) { t.Parallel() + type logLine struct { + level log.Level + message string + } + testCases := map[string]struct { b []byte hardcodedVersions map[string]uint16 - logged []string + logged []logLine persisted models.AllServers errMessage string }{ @@ -42,7 +48,9 @@ func Test_extractServersFromBytes(t *testing.T) { b: []byte(`{"cyberghost": "garbage"}`), hardcodedVersions: populateProviderToVersion(map[string]uint16{}), 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 int64 \"json:\\\"timestamp\\\"\"; " + + "Filepath string \"json:\\\"filepath\\\"\" }", }, "bad servers array JSON": { b: []byte(`{"cyberghost": {"version": 1, "servers": "garbage"}}`), @@ -81,13 +89,30 @@ func Test_extractServersFromBytes(t *testing.T) { hardcodedVersions: populateProviderToVersion(map[string]uint16{ providers.Cyberghost: 2, }), - logged: []string{ - "Cyberghost servers from file discarded because they have version 1 and hardcoded servers have version 2", + logged: []logLine{ + {level: log.LevelInfo, message: "Cyberghost servers from file discarded because they have version 1" + + " and hardcoded servers have version 2"}, }, 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, + }), + logged: []logLine{ + {level: log.LevelWarn, message: "Cyberghost preferred servers from file discarded because they have version 1" + + " and hardcoded servers have version 2"}, + }, + persisted: models.AllServers{ + ProviderToServers: map[string]models.Servers{}, + }, + errMessage: "", + }, } for name, testCase := range testCases { @@ -98,7 +123,15 @@ func Test_extractServersFromBytes(t *testing.T) { logger := NewMockLogger(ctrl) var previousLogCall *gomock.Call for _, logged := range testCase.logged { - call := logger.EXPECT().Info(logged) + var call *gomock.Call + switch logged.level { //nolint:exhaustive + case log.LevelInfo: + call = logger.EXPECT().Info(logged.message) + case log.LevelWarn: + call = logger.EXPECT().Warn(logged.message) + default: + t.Fatalf("invalid log level %d in test case", logged.level) + } if previousLogCall != nil { call.After(previousLogCall) } diff --git a/internal/storage/servers.go b/internal/storage/servers.go index d02eeebb..111ffe64 100644 --- a/internal/storage/servers.go +++ b/internal/storage/servers.go @@ -2,6 +2,8 @@ package storage import ( "fmt" + "os" + "path/filepath" "time" "github.com/qdm12/gluetun/internal/constants/providers" @@ -10,12 +12,12 @@ import ( // SetServers sets the given servers for the given provider // 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 // NOT MUTATE them after calling this method. func (s *Storage) SetServers(provider string, servers []models.Server) (err error) { if provider == providers.Custom { - return + return nil } s.mergedMutex.Lock() @@ -26,10 +28,24 @@ func (s *Storage) SetServers(provider string, servers []models.Server) (err erro serversObject.Servers = servers s.mergedServers.ProviderToServers[provider] = serversObject - err = s.flushToFile(s.filepath) + if s.directoryPath == "" { + return nil // no disk writing + } + + manifestPath := filepath.Join(s.directoryPath, manifestFilename) + err = s.flushToFile(manifestPath) if err != nil { 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 && !os.IsNotExist(err) { + s.logger.Warn("failed removing legacy servers file " + s.legacyFilepath + ": " + err.Error()) + } return nil } diff --git a/internal/storage/storage.go b/internal/storage/storage.go index ab0378f8..88366e11 100644 --- a/internal/storage/storage.go +++ b/internal/storage/storage.go @@ -1,6 +1,8 @@ package storage import ( + "os" + "path/filepath" "sync" "github.com/qdm12/gluetun/internal/models" @@ -14,31 +16,36 @@ type Storage struct { // SyncServers method. hardcodedServers models.AllServers logger Logger - filepath string + directoryPath string + legacyFilepath string } +const manifestFilename = "manifest.json" + type Logger interface { Info(s string) + Infof(format string, args ...any) Warn(s string) } // New creates a new storage and reads the servers from the -// embedded servers file and the file on disk. -// Passing an empty filepath disables the reading and writing of +// embedded servers files and the files on disk. +// Passing an empty directoryPath disables the reading and writing of // servers. -func New(logger Logger, filepath string) (storage *Storage, err error) { - // A unit test prevents any error from being returned +func New(logger Logger, directoryPath, legacyFilepath string) (storage *Storage, err error) { + // A unit test prevents [parseHardcodedServers] from ever failing, // and ensures all providers are part of the servers returned. - hardcodedServers, _ := parseHardcodedServers() + hardcodedServers := parseHardcodedServers() storage = &Storage{ hardcodedServers: hardcodedServers, mergedServers: hardcodedServers, logger: logger, - filepath: filepath, + directoryPath: directoryPath, + legacyFilepath: legacyFilepath, } - if filepath != "" { + if directoryPath != "" { if err := storage.syncServers(); err != nil { return nil, err } @@ -46,3 +53,18 @@ func New(logger Logger, filepath string) (storage *Storage, err error) { 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.legacyFilepath == "" { + 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() +} diff --git a/internal/storage/sync.go b/internal/storage/sync.go index de355c9c..540a9895 100644 --- a/internal/storage/sync.go +++ b/internal/storage/sync.go @@ -2,6 +2,8 @@ package storage import ( "fmt" + "os" + "path/filepath" "reflect" "github.com/qdm12/gluetun/internal/models" @@ -14,18 +16,31 @@ func countServers(allServers models.AllServers) (count int) { 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) { hardcodedVersions := make(map[string]uint16, len(s.hardcodedServers.ProviderToServers)) for provider, servers := range s.hardcodedServers.ProviderToServers { 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 { 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) countOnFile := countServers(serversOnFile) @@ -34,13 +49,13 @@ func (s *Storage) syncServers() (err error) { if countOnFile == 0 { s.logger.Info(fmt.Sprintf( - "creating %s with %d hardcoded servers", - s.filepath, hardcodedCount)) + "writing servers data files to %s with %d hardcoded servers", + s.directoryPath, hardcodedCount)) s.mergedServers = s.hardcodedServers } else { s.logger.Info(fmt.Sprintf( - "merging by most recent %d hardcoded servers and %d servers read from %s", - hardcodedCount, countOnFile, s.filepath)) + "merging by most recent %d hardcoded servers and %d servers read from manifest file %s", + hardcodedCount, countOnFile, sourceManifestPath)) s.mergedServers = s.mergeServers(s.hardcodedServers, serversOnFile) } @@ -50,9 +65,18 @@ func (s *Storage) syncServers() (err error) { return nil } - err = s.flushToFile(s.filepath) + err = s.flushToFile(destinationManifestPath) 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 }