to re-review

This commit is contained in:
Quentin McGaw
2026-05-12 04:43:35 +00:00
parent 71c4fcda52
commit aa6b2e1a4d
5 changed files with 193 additions and 46 deletions
+73 -21
View File
@@ -12,29 +12,30 @@ 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) 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 +47,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 +93,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 int64 `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 +115,55 @@ func (s *Storage) readServers(provider string, hardcodedVersion uint16,
provider, err) 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
} }
+38 -5
View File
@@ -7,6 +7,7 @@ import (
"github.com/golang/mock/gomock" "github.com/golang/mock/gomock"
"github.com/qdm12/gluetun/internal/constants/providers" "github.com/qdm12/gluetun/internal/constants/providers"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
"github.com/qdm12/log"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
@@ -27,10 +28,15 @@ func populateProviderToVersion(providerToVersion map[string]uint16) map[string]u
func Test_extractServersFromBytes(t *testing.T) { func Test_extractServersFromBytes(t *testing.T) {
t.Parallel() t.Parallel()
type logLine struct {
level log.Level
message string
}
testCases := map[string]struct { testCases := map[string]struct {
b []byte b []byte
hardcodedVersions map[string]uint16 hardcodedVersions map[string]uint16
logged []string logged []logLine
persisted models.AllServers persisted models.AllServers
errMessage string errMessage string
}{ }{
@@ -42,7 +48,9 @@ func Test_extractServersFromBytes(t *testing.T) {
b: []byte(`{"cyberghost": "garbage"}`), b: []byte(`{"cyberghost": "garbage"}`),
hardcodedVersions: populateProviderToVersion(map[string]uint16{}), hardcodedVersions: populateProviderToVersion(map[string]uint16{}),
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 int64 \"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"}}`),
@@ -81,13 +89,30 @@ func Test_extractServersFromBytes(t *testing.T) {
hardcodedVersions: populateProviderToVersion(map[string]uint16{ hardcodedVersions: populateProviderToVersion(map[string]uint16{
providers.Cyberghost: 2, providers.Cyberghost: 2,
}), }),
logged: []string{ logged: []logLine{
"Cyberghost servers from file discarded because they have version 1 and hardcoded servers have version 2", {level: log.LevelInfo, message: "Cyberghost servers from file discarded because they have version 1" +
" and hardcoded servers have version 2"},
}, },
persisted: models.AllServers{ persisted: models.AllServers{
ProviderToServers: map[string]models.Servers{}, 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 { for name, testCase := range testCases {
@@ -98,7 +123,15 @@ func Test_extractServersFromBytes(t *testing.T) {
logger := NewMockLogger(ctrl) logger := NewMockLogger(ctrl)
var previousLogCall *gomock.Call var previousLogCall *gomock.Call
for _, logged := range testCase.logged { 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 { if previousLogCall != nil {
call.After(previousLogCall) call.After(previousLogCall)
} }
+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.directoryPath == "" {
return nil // no disk writing
}
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 && !os.IsNotExist(err) {
s.logger.Warn("failed removing legacy servers file " + s.legacyFilepath + ": " + err.Error())
}
return nil return nil
} }
+30 -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,36 @@ type Storage struct {
// SyncServers method. // SyncServers method.
hardcodedServers models.AllServers hardcodedServers models.AllServers
logger Logger logger Logger
filepath string 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, 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, directoryPath: directoryPath,
legacyFilepath: legacyFilepath,
} }
if filepath != "" { if directoryPath != "" {
if err := storage.syncServers(); err != nil { if err := storage.syncServers(); err != nil {
return nil, err return nil, err
} }
@@ -46,3 +53,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.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()
}
+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
} }