From 0eea70067b25c7955c167da972c3227f83a2f711 Mon Sep 17 00:00:00 2001 From: Jay Brown Date: Wed, 5 Feb 2025 20:03:36 +0000 Subject: [PATCH] Merged in bugfix/baseconfig (pull request #46) Fix bug and add unit test to verify * fix bug and test * pr comments * move to server * fix tests --- internal/server/server.go | 9 +++ internal/server/service/listener.go | 1 + internal/serviceconfig/common.go | 80 ++++++++++++++++++++ internal/serviceconfig/common_test.go | 61 +++++++++++++++ internal/serviceconfig/logger/config.go | 64 ---------------- internal/serviceconfig/logger/config_test.go | 21 ----- 6 files changed, 151 insertions(+), 85 deletions(-) diff --git a/internal/server/server.go b/internal/server/server.go index b4a37878..6040e1c6 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -8,6 +8,8 @@ import ( "github.com/go-playground/validator/v10" ) +const DEFAULT_SECRET_PREFIX = "secret" + type Config interface { serviceconfig.ConfigProvider SetValidator() @@ -29,6 +31,13 @@ func (c *BaseConfig) GetValidator() *validator.Validate { } func New(ctx context.Context, cfg Config) (func() error, error) { + // init the config in case it has not already been called. + errInitializingConfig := serviceconfig.InitializeConfig(cfg) + if errInitializingConfig != nil { + return func() error { return nil }, errInitializingConfig + } + + cfg.PrintConfig(DEFAULT_SECRET_PREFIX) closeTracer := cfg.SetOtel(ctx) err := migrations.Run(ctx, cfg) diff --git a/internal/server/service/listener.go b/internal/server/service/listener.go index 3f931acb..c5014135 100644 --- a/internal/server/service/listener.go +++ b/internal/server/service/listener.go @@ -60,6 +60,7 @@ type Server struct { } func New(ctx context.Context, cfg Config) (*Server, error) { + cleanup, err := server.New(ctx, cfg) if err != nil { return nil, err diff --git a/internal/serviceconfig/common.go b/internal/serviceconfig/common.go index f7ea73b2..a4ae98c8 100644 --- a/internal/serviceconfig/common.go +++ b/internal/serviceconfig/common.go @@ -2,6 +2,7 @@ package serviceconfig import ( "errors" + "fmt" "log/slog" "os" "queryorchestration/internal/serviceconfig/aws" @@ -10,6 +11,7 @@ import ( "queryorchestration/internal/serviceconfig/observability" "queryorchestration/internal/serviceconfig/queue" "reflect" + "strings" "github.com/caarlos0/env/v11" "github.com/joho/godotenv" @@ -136,3 +138,81 @@ func (b *BaseConfig) SetDBConfig(cfg *database.DBConfig) { func (b *BaseConfig) SetAWSConfig(cfg *aws.AWSConfig) { b.AWSConfig = *cfg } + +func (b *BaseConfig) PrintConfig(prefixSecret string) { + b.printConfigRecursive(reflect.ValueOf(b), "", prefixSecret, make(map[uintptr]bool)) +} + +func (b *BaseConfig) printConfigRecursive(val reflect.Value, prefix string, prefixSecret string, visited map[uintptr]bool) { + // Handle pointer dereference + if val.Kind() == reflect.Ptr { + val = val.Elem() + } + + // Prevent infinite recursion using the memory address + // Recommended way to get the uintptr + addr := uintptr(val.Addr().UnsafePointer()) + if visited[addr] { + return + } + visited[addr] = true + + typ := val.Type() + + for i := 0; i < val.NumField(); i++ { + field := val.Field(i) + fieldType := typ.Field(i) + + // Skip unexported fields + if !fieldType.IsExported() { + continue + } + + fieldName := fieldType.Name + fullPath := prefix + fieldName + + // Check for env tag + envTag := fieldType.Tag.Get("env") + if envTag == "" { + // If this is a struct, we still need to recurse into it + // as it might contain fields with env tags + if field.Kind() == reflect.Struct { + b.printConfigRecursive(field, fullPath+".", prefixSecret, visited) + } + continue + } + + // Handle different kinds of fields + switch field.Kind() { + case reflect.Struct: + // If it's a struct AND has an env tag, print it and recurse + var valueStr string + if field.CanInterface() { + valueStr = fmt.Sprintf("%+v", field.Interface()) + } + + b.Logger.Info("Struct Config value", + "key", fullPath, + "value", valueStr) + b.printConfigRecursive(field, fullPath+".", prefixSecret, visited) + default: + var valueStr string + if field.Kind() == reflect.String { + valueStr = field.String() + } else { + valueStr = fmt.Sprintf("%v", field.Interface()) + } + + // Mask sensitive values + if strings.Contains(strings.ToLower(fieldName), strings.ToLower(prefixSecret)) { + if len(valueStr) > 3 { + valueStr = valueStr[:3] + "..." + } + } + + b.Logger.Info("Config value", + "key", fullPath, + "value", valueStr) + } + } +} diff --git a/internal/serviceconfig/common_test.go b/internal/serviceconfig/common_test.go index bf8eb7cb..24eac97c 100644 --- a/internal/serviceconfig/common_test.go +++ b/internal/serviceconfig/common_test.go @@ -1,6 +1,8 @@ package serviceconfig import ( + "bytes" + "log/slog" "os" "queryorchestration/internal/serviceconfig/aws" "queryorchestration/internal/serviceconfig/database" @@ -174,3 +176,62 @@ func TestSetAWSConfig(t *testing.T) { cfg.SetAWSConfig(&acfg) assert.Equal(t, acfg, cfg.AWSConfig) } + +func TestPrintConfigRecursive(t *testing.T) { + envVars := map[string]string{ + "DB_USER": "postgres", + "DB_PASS": "password", + "DB_HOST": "localhost", + "DB_PORT": "5432", + "DB_NAME": "query_orchestration", + "DB_NOSSL": "true", + "AWS_ACCESS_KEY_ID": "key", + "AWS_SECRET_ACCESS_KEY": "password", + "AWS_REGION": "region", + "PWD": "/foo", + } + + // Set environment variables + for k, v := range envVars { + t.Setenv(k, v) + } + + // Create a buffer to capture log output + var logBuffer bytes.Buffer + + // Create a test logger that writes to our buffer + testLogger := slog.New(slog.NewTextHandler(&logBuffer, &slog.HandlerOptions{ + Level: slog.LevelInfo, + })) + + cfg := &BaseConfig{} + err := InitializeConfig(cfg) + assert.NoError(t, err) + + cfg.Logger = testLogger + + // Call the function we're testing + cfg.PrintConfig("secret") + + // Get the log output as a string + logOutput := logBuffer.String() + + // Check that each environment variable value appears in the logs + // except for DB_PASS which should be masked + for envKey, envValue := range envVars { + if envKey == "DB_PASS" || envKey == "AWS_SECRET_ACCESS_KEY" { + // Password should be masked + if !strings.Contains(logOutput, "pas") { + t.Errorf("Expected masked password in logs, got: %s", logOutput) + } + // Full password should not appear + if strings.Contains(logOutput, envValue) { + t.Errorf("Full password should not appear in logs: %s", logOutput) + } + } else { + if !strings.Contains(logOutput, envValue) { + t.Errorf("Expected %s in logs, but it was not found. Log output: %s", envValue, logOutput) + } + } + } +} diff --git a/internal/serviceconfig/logger/config.go b/internal/serviceconfig/logger/config.go index 35d78f4f..663542e4 100644 --- a/internal/serviceconfig/logger/config.go +++ b/internal/serviceconfig/logger/config.go @@ -1,10 +1,7 @@ package logger import ( - "fmt" "log/slog" - "reflect" - "strings" ) type LogConfig struct { @@ -19,64 +16,3 @@ type ConfigProvider interface { func (b *LogConfig) GetLogger() *slog.Logger { return b.Logger } - -func (b *LogConfig) PrintConfig(prefixSecret string) { - b.printConfigRecursive(reflect.ValueOf(b), "", prefixSecret, make(map[reflect.Value]bool)) -} - -func (b *LogConfig) printConfigRecursive(val reflect.Value, prefix string, prefixSecret string, visited map[reflect.Value]bool) { - // Handle pointer dereference - if val.Kind() == reflect.Ptr { - val = val.Elem() - } - - // Prevent infinite recursion - if visited[val] { - return - } - visited[val] = true - - typ := val.Type() - - for i := 0; i < val.NumField(); i++ { - field := val.Field(i) - fieldType := typ.Field(i) - - // Skip unexported fields - if !fieldType.IsExported() { - continue - } - - fieldName := fieldType.Name - fullPath := prefix + fieldName - - // Handle embedded fields - if fieldType.Anonymous { - b.printConfigRecursive(field, prefix, prefixSecret, visited) - continue - } - - switch field.Kind() { - case reflect.Struct: - b.printConfigRecursive(field, fullPath+".", prefixSecret, visited) - default: - var valueStr string - if field.Kind() == reflect.String { - valueStr = field.String() - } else { - valueStr = fmt.Sprintf("%v", field.Interface()) - } - - // Mask sensitive values - if strings.Contains(strings.ToLower(fieldName), strings.ToLower(prefixSecret)) { - if len(valueStr) > 5 { - valueStr = valueStr[:5] + "..." - } - } - - b.Logger.Info("Config value", - "key", fullPath, - "value", valueStr) - } - } -} diff --git a/internal/serviceconfig/logger/config_test.go b/internal/serviceconfig/logger/config_test.go index 6d623b2a..89727273 100644 --- a/internal/serviceconfig/logger/config_test.go +++ b/internal/serviceconfig/logger/config_test.go @@ -2,7 +2,6 @@ package logger import ( "log/slog" - "reflect" "testing" "github.com/stretchr/testify/assert" @@ -14,23 +13,3 @@ func TestGetLogger(t *testing.T) { cfg.Logger = slog.Default() assert.NotNil(t, cfg.GetLogger()) } - -func TestPrintConfig(t *testing.T) { - cfg := &LogConfig{} - tl := &TestLogger{T: t} - cfg.Logger = slog.New(tl) - cfg.PrintConfig("") - assert.Len(t, tl.Logs, 1) - assert.Equal(t, "Logger", tl.Logs[0]["key"]) -} - -func TestPrintConfigRecursive(t *testing.T) { - cfg := &LogConfig{} - tl := &TestLogger{T: t} - cfg.Logger = slog.New(tl) - v := struct{ Example string }{Example: "examplestring"} - cfg.printConfigRecursive(reflect.ValueOf(v), "", "", make(map[reflect.Value]bool)) - assert.Len(t, tl.Logs, 1) - assert.Equal(t, "Example", tl.Logs[0]["key"]) - assert.Equal(t, "examp...", tl.Logs[0]["value"]) -}