From ad8ed5130c18a265e58c5d99d659206196260182 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Sun, 11 Aug 2024 22:18:49 +0300 Subject: add CBs for config defaults --- internal/config/client/client.go | 12 ++++----- internal/config/config.go | 25 ++++++++++++------- internal/config/config_test.go | 54 ++++++++++++++++++++++++++++++++-------- internal/config/server/server.go | 24 ++++++++---------- 4 files changed, 76 insertions(+), 39 deletions(-) diff --git a/internal/config/client/client.go b/internal/config/client/client.go index 8676114..00819c0 100644 --- a/internal/config/client/client.go +++ b/internal/config/client/client.go @@ -21,16 +21,16 @@ type ClientConfig struct { func New(configFile string) (ClientConfig, error) { conf, _ := config.FromFile[ClientConfig](configFile) - conf.Server = config.FromENV("GOS_SERVERS", conf.Server) - conf.APIKey = config.FromENV("GOS_API_KEY", conf.APIKey) - conf.Editor = config.FromENV("GOS_EDITOR", "EDITOR", conf.Editor, "vi") + conf.Server = config.EnvToStr("GOS_SERVERS", conf.Server) + conf.APIKey = config.EnvToStr("GOS_API_KEY", conf.APIKey) + conf.Editor = config.EnvToStr("GOS_EDITOR", "EDITOR", conf.Editor, "vi") defaultDataDir := fmt.Sprintf("%s/.gos/data", os.Getenv("HOME")) - conf.DataDir = config.FromENV("GOS_DATA_DIR", conf.DataDir, defaultDataDir) - conf.ComposeFile = config.FromENV("GOS_COMPOSE_FILE", conf.ComposeFile, "compose.txt") + conf.DataDir = config.EnvToStr("GOS_DATA_DIR", conf.DataDir, defaultDataDir) + conf.ComposeFile = config.EnvToStr("GOS_COMPOSE_FILE", conf.ComposeFile, "compose.txt") defaultLogFile := fmt.Sprintf("%s/.gos/gos.log", os.Getenv("HOME")) - conf.LogFile = config.FromENV("GOS_LOG_FILE", conf.LogFile, defaultLogFile) + conf.LogFile = config.EnvToStr("GOS_LOG_FILE", conf.LogFile, defaultLogFile) return conf, nil } diff --git a/internal/config/config.go b/internal/config/config.go index 1ebfc68..12d40f0 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -27,16 +27,21 @@ func FromFile[T any](configFile string) (T, error) { } // Set config from envoronment variable if present, e.g. hansWurst from GOS_HANS_WURST -func EnvToStr(keys ...string) string { +func EnvToStr(keys ...any) string { for _, key := range keys { - if key == "" { - continue - } - if !isAllUpperCase(key) { - return key - } - if value := os.Getenv(key); value != "" { - return value + switch key := key.(type) { + case string: + if key == "" { + continue + } + if !isAllUpperCase(key) { + return key + } + if value := os.Getenv(key); value != "" { + return value + } + case func() string: + return key() } } @@ -59,6 +64,8 @@ func EnvToInt(keys ...any) int { } case int: return key + case func() int: + return key() } } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index e1329d3..27da6ed 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -5,14 +5,14 @@ import ( "testing" ) -func TestFromENV(t *testing.T) { +func TestEnvToStr(t *testing.T) { t.Parallel() os.Setenv("GOS_TEST_FROM_ENV", "foobarbaz") var ( expected = "foobarbaz" - got = FromENV("GOS_TEST_FROM_ENV") + got = EnvToStr("GOS_TEST_FROM_ENV") ) if got != expected { @@ -21,34 +21,34 @@ func TestFromENV(t *testing.T) { t.Logf("got '%s' as expected", expected) expected = "default value" - got = FromENV("GOS_JAJAJA", expected) + got = EnvToStr("GOS_JAJAJA", expected) if got != expected { t.Errorf("got '%s' but expected '%s'", got, expected) } t.Logf("got '%s' as expected", expected) os.Unsetenv("JUJUJU_NOT_EXISTANT_ENV") - if got = FromENV("JUJUJU_NOT_EXISTANT_ENV"); got != "" { + if got = EnvToStr("JUJUJU_NOT_EXISTANT_ENV"); got != "" { t.Errorf("got '%s' but expected empty string", got) } t.Logf("got empty string as expected") expected = "casio g-shock" - got = FromENV("GOS_WATCH", "", "", "", expected, "") + got = EnvToStr("GOS_WATCH", "", "", "", expected, "") if got != expected { t.Errorf("got '%s' but expected '%s'", got, expected) } t.Logf("got '%s' as expected", expected) } -func TestIntFromENV(t *testing.T) { +func TestEnvToInt(t *testing.T) { t.Parallel() os.Setenv("GOS_TEST_INT_FROM_ENV", "1") var ( expected = 1 - got = IntFromENV(t, "GOS_TEST_INT_FROM_ENV") + got = EnvToInt(t, "GOS_TEST_INT_FROM_ENV") ) if got != expected { @@ -57,20 +57,20 @@ func TestIntFromENV(t *testing.T) { t.Logf("got '%d' as expected", expected) expected = 999 - got = IntFromENV("GOS_JAJAJA", expected) + got = EnvToInt("GOS_JAJAJA", expected) if got != expected { t.Errorf("got '%d' but expected '%d'", got, expected) } t.Logf("got '%d' as expected", expected) os.Unsetenv("JUJUJU_NOT_EXISTANT_ENV") - if got = IntFromENV("JUJUJU_NOT_EXISTANT_ENV"); got != 0 { + if got = EnvToInt("JUJUJU_NOT_EXISTANT_ENV"); got != 0 { t.Errorf("got '%d' but expected zero", got) } t.Logf("got zero as expected") expected = 1234 - got = IntFromENV("GOS_WATCH", "", "", "", expected, "") + got = EnvToInt("GOS_WATCH", "", "", "", expected, "") if got != expected { t.Errorf("got '%d' but expected '%d'", got, expected) } @@ -85,7 +85,7 @@ func TestSecondENV(t *testing.T) { var ( expected = "hx" - got = FromENV("GOS_NONEXISTANT", "EDITOR", "notepad.exe") + got = EnvToStr("GOS_NONEXISTANT", "EDITOR", "notepad.exe") ) if expected != got { @@ -104,3 +104,35 @@ func TestIsAllUpperCase(t *testing.T) { t.Errorf("should be all upper") } } + +func TestDefaultStrCB(t *testing.T) { + t.Parallel() + os.Unsetenv("GOS_NONEXISTANT") + + var ( + expected = "hello" + got = EnvToStr("GOS_NONEXISTANT", func() string { + return "hello" + }) + ) + + if expected != got { + t.Errorf("got '%s' but expected '%s'", got, expected) + } +} + +func TestDefaultIntCB(t *testing.T) { + t.Parallel() + os.Unsetenv("GOS_NONEXISTANT") + + var ( + expected = 666 + got = EnvToInt("GOS_NONEXISTANT", func() int { + return 666 + }) + ) + + if expected != got { + t.Errorf("got '%d' but expected '%d'", got, expected) + } +} diff --git a/internal/config/server/server.go b/internal/config/server/server.go index d46a32a..a123e81 100644 --- a/internal/config/server/server.go +++ b/internal/config/server/server.go @@ -23,24 +23,22 @@ type ServerConfig struct { func New(configFile string) (ServerConfig, error) { conf, _ := config.FromFile[ServerConfig](configFile) - conf.ListenAddr = config.StrFromENV("GOS_LISTEN_ADDR", conf.ListenAddr, "localhost:8080") - conf.Partner = config.StrFromENV("GOS_PARTNER", conf.Partner) - conf.APIKey = config.StrFromENV("GOS_API_KEY", conf.APIKey) - conf.DataDir = config.StrFromENV("GOS_DATA_DIR", conf.DataDir, "data") - conf.EmailTo = config.StrFromENV("GOS_EMAIL_TO", conf.EmailTo) - conf.EmailFrom = config.StrFromENV("GOS_EMAIL_FROM", conf.EmailFrom) - - conf.SMTPServer = config.StrFromENV("GOS_SMTP_SERVER", conf.SMTPServer) - if conf.SMTPServer == "" { + conf.ListenAddr = config.EnvToStr("GOS_LISTEN_ADDR", conf.ListenAddr, "localhost:8080") + conf.Partner = config.EnvToStr("GOS_PARTNER", conf.Partner) + conf.APIKey = config.EnvToStr("GOS_API_KEY", conf.APIKey) + conf.DataDir = config.EnvToStr("GOS_DATA_DIR", conf.DataDir, "data") + conf.EmailTo = config.EnvToStr("GOS_EMAIL_TO", conf.EmailTo) + conf.EmailFrom = config.EnvToStr("GOS_EMAIL_FROM", conf.EmailFrom) + + conf.SMTPServer = config.EnvToStr("GOS_SMTP_SERVER", conf.SMTPServer, func() string { hostname, err := os.Hostname() if err != nil { log.Fatal(err) } - conf.SMTPServer = fmt.Sprintf("%s:25", hostname) - log.Println("Set SMTPServer to " + conf.SMTPServer) - } + return fmt.Sprintf("%s:25", hostname) + }) - conf.CRONMergeIntervalS = config.IntFromENV("GOS_CRON_MERGE_INTERVAL", 3600) + conf.CRONMergeIntervalS = config.EnvToInt("GOS_CRON_MERGE_INTERVAL", 3600) return conf, nil } -- cgit v1.2.3