mirror of
https://github.com/qdm12/gluetun.git
synced 2026-07-22 10:26:26 +02:00
to re-review
This commit is contained in:
+73
-21
@@ -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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user