diff --git a/management/cmd/config.go b/management/cmd/config.go index b41c8addb..bc8465322 100644 --- a/management/cmd/config.go +++ b/management/cmd/config.go @@ -3,11 +3,19 @@ package cmd import ( nbconfig "github.com/netbirdio/netbird/management/internals/server/config" configloader "github.com/netbirdio/netbird/util/config" + "github.com/netbirdio/netbird/util/envtemplate" ) func loadManagementConfig(configPath string) (*nbconfig.Config, error) { - return configloader.Load(configPath, &nbconfig.Config{Datadir: defaultMgmtDataDir}, configloader.Options{ + cfg, err := configloader.Load(configPath, &nbconfig.Config{Datadir: defaultMgmtDataDir}, configloader.Options{ TagName: "json", - Transform: configloader.ExpandEnvTemplate, + Transform: envtemplate.Expand, }) + if err != nil { + return nil, err + } + if cfg.Datadir == "" { + cfg.Datadir = defaultMgmtDataDir + } + return cfg, nil } diff --git a/management/cmd/config_test.go b/management/cmd/config_test.go index 9a78c5a00..009ba702f 100644 --- a/management/cmd/config_test.go +++ b/management/cmd/config_test.go @@ -10,6 +10,15 @@ import ( "github.com/stretchr/testify/require" ) +func TestLoadManagementConfigUsesDefaultDataDirForEmptyValue(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "management.json") + require.NoError(t, os.WriteFile(configPath, []byte(`{"Datadir":""}`), 0o600)) + + cfg, err := loadManagementConfig(configPath) + require.NoError(t, err) + assert.Equal(t, defaultMgmtDataDir, cfg.Datadir, "Empty legacy values should use the default data directory") +} + func TestLoadManagementConfigSources(t *testing.T) { t.Setenv("MANAGEMENT_DATA_DIR", "/template-data") t.Setenv("MANAGEMENT_ENCRYPTION_KEY", "template-key") diff --git a/util/config/loader.go b/util/config/loader.go index 69801a375..fc56f098d 100644 --- a/util/config/loader.go +++ b/util/config/loader.go @@ -28,8 +28,7 @@ // }) // // Set [Options.AllowMissing] when the service must start without a configuration -// file. [Options.Transform] can preprocess file contents before decoding, such -// as with [ExpandEnvTemplate]. +// file. [Options.Transform] can preprocess file contents before decoding. package config import ( @@ -215,66 +214,112 @@ func bindConfigSources( defer delete(visiting, configType) for i := range configType.NumField() { - field := configType.Field(i) - if !field.IsExported() { - continue + if err := bindConfigField( + v, + configType.Field(i), + prefix, + tagName, + flagSet, + bindEnvironment, + bindFlags, + visiting, + ); err != nil { + return err } + } + return nil +} - key, inline, skip := configFieldKey(field, tagName) - if skip { - continue - } - if inline { - key = prefix - } else if prefix != "" { - key = prefix + "." + key - } +func bindConfigField( + v *viper.Viper, + field reflect.StructField, + prefix string, + tagName string, + flagSet *pflag.FlagSet, + bindEnvironment bool, + bindFlags bool, + visiting map[reflect.Type]bool, +) error { + if !field.IsExported() { + return nil + } - fieldEnvironment := field.Tag.Get("env") - fieldFlags := field.Tag.Get("flag") - bindFieldEnvironment := bindEnvironment && fieldEnvironment != "-" - bindFieldFlags := bindFlags && fieldFlags != "-" + key, inline, skip := configFieldKey(field, tagName) + if skip { + return nil + } + if inline { + key = prefix + } else if prefix != "" { + key = prefix + "." + key + } - fieldType := field.Type - for fieldType.Kind() == reflect.Pointer { - fieldType = fieldType.Elem() - } - if fieldType.Kind() == reflect.Struct && !isScalarUnmarshaler(fieldType) { - if !visiting[fieldType] { - if err := bindConfigSources( - v, - fieldType, - key, - tagName, - flagSet, - bindFieldEnvironment, - bindFieldFlags, - visiting, - ); err != nil { - return err - } - } - continue - } + fieldEnvironment := field.Tag.Get("env") + fieldFlags := field.Tag.Get("flag") + bindEnvironment = bindEnvironment && fieldEnvironment != "-" + bindFlags = bindFlags && fieldFlags != "-" - if key == "" { - return fmt.Errorf("empty config key for field %s", field.Name) + fieldType := field.Type + for fieldType.Kind() == reflect.Pointer { + fieldType = fieldType.Elem() + } + if fieldType.Kind() == reflect.Struct && !isScalarUnmarshaler(fieldType) { + if visiting[fieldType] { + return nil } - if bindFieldEnvironment { - if err := bindEnvironmentVariable(v, key, fieldEnvironment); err != nil { - return err - } - } - if bindFieldFlags && flagSet != nil && fieldFlags != "" { - flagName, flag, err := selectFlag(flagSet, fieldFlags) - if err != nil { - return fmt.Errorf("config field %s: %w", field.Name, err) - } - if err := v.BindPFlag(key, flag); err != nil { - return fmt.Errorf("bind flag %s: %w", flagName, err) - } + return bindConfigSources( + v, + fieldType, + key, + tagName, + flagSet, + bindEnvironment, + bindFlags, + visiting, + ) + } + + return bindScalarSources( + v, + field, + key, + fieldEnvironment, + fieldFlags, + flagSet, + bindEnvironment, + bindFlags, + ) +} + +func bindScalarSources( + v *viper.Viper, + field reflect.StructField, + key string, + environmentName string, + flagNames string, + flagSet *pflag.FlagSet, + bindEnvironment bool, + bindFlags bool, +) error { + if key == "" { + return fmt.Errorf("empty config key for field %s", field.Name) + } + if bindEnvironment { + if err := bindEnvironmentVariable(v, key, environmentName); err != nil { + return err } } + if !bindFlags || flagSet == nil || flagNames == "" { + return nil + } + + flagName, flag, err := selectFlag(flagSet, flagNames) + if err != nil { + return fmt.Errorf("config field %s: %w", field.Name, err) + } + if err := v.BindPFlag(key, flag); err != nil { + return fmt.Errorf("bind flag %s: %w", flagName, err) + } return nil } diff --git a/util/config/loader_test.go b/util/config/loader_test.go index 82da27c48..9dcb602b8 100644 --- a/util/config/loader_test.go +++ b/util/config/loader_test.go @@ -11,6 +11,8 @@ import ( "github.com/spf13/pflag" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/util/envtemplate" ) type testConfig struct { @@ -138,7 +140,7 @@ server: cfg, err := Load(configPath, defaultTestConfig(), Options{ TagName: "yaml", - Transform: ExpandEnvTemplate, + Transform: envtemplate.Expand, }) require.NoError(t, err) assert.Equal(t, ":8443", cfg.Server.Address, "The transform should run before decoding") diff --git a/util/config/template.go b/util/envtemplate/template.go similarity index 53% rename from util/config/template.go rename to util/envtemplate/template.go index 618c3d6f5..753b89a09 100644 --- a/util/config/template.go +++ b/util/envtemplate/template.go @@ -1,4 +1,5 @@ -package config +// Package envtemplate expands Go templates with environment variables. +package envtemplate import ( "bytes" @@ -8,27 +9,27 @@ import ( "text/template" ) -// ExpandEnvTemplate substitutes Go-template references with environment values. -func ExpandEnvTemplate(data []byte) ([]byte, error) { +// Expand substitutes Go-template references with environment values. +func Expand(data []byte) ([]byte, error) { tmpl, err := template.New("config").Parse(string(data)) if err != nil { return nil, fmt.Errorf("parse environment template: %w", err) } var output bytes.Buffer - if err := tmpl.Execute(&output, environmentMap()); err != nil { + if err := tmpl.Execute(&output, environment()); err != nil { return nil, fmt.Errorf("execute environment template: %w", err) } return output.Bytes(), nil } -func environmentMap() map[string]string { - environment := make(map[string]string) +func environment() map[string]string { + values := make(map[string]string) for _, entry := range os.Environ() { key, value, ok := strings.Cut(entry, "=") if ok { - environment[key] = value + values[key] = value } } - return environment + return values } diff --git a/util/file.go b/util/file.go index 7d3944bf2..8254368a3 100644 --- a/util/file.go +++ b/util/file.go @@ -12,7 +12,7 @@ import ( log "github.com/sirupsen/logrus" - configloader "github.com/netbirdio/netbird/util/config" + "github.com/netbirdio/netbird/util/envtemplate" ) func WriteBytesWithRestrictedPermission(ctx context.Context, file string, bs []byte) error { @@ -243,7 +243,7 @@ func ReadJsonWithEnvSub(file string, res interface{}) (interface{}, error) { return nil, err } - output, err := configloader.ExpandEnvTemplate(bs) + output, err := envtemplate.Expand(bs) if err != nil { return nil, err }