Files
query-orchestration/internal/serviceconfig/common_test.go
T
Jay Brown 0eea70067b 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
2025-02-05 20:03:36 +00:00

238 lines
6.3 KiB
Go

package serviceconfig
import (
"bytes"
"log/slog"
"os"
"queryorchestration/internal/serviceconfig/aws"
"queryorchestration/internal/serviceconfig/database"
"strings"
"testing"
"github.com/stretchr/testify/assert"
)
type TestingServiceNameConfig struct {
BaseConfig
AppEnv string `env:"APP_ENV"`
BoolTest bool `env:"BOOL_TEST"`
IntTest int `env:"INT_TEST,required,notEmpty"`
}
func TestInitializeConfig(t *testing.T) {
tests := []struct {
name string
envVars map[string]string
wantErr bool
errMessageContains string
}{
{
name: "valid configuration",
envVars: map[string]string{
"APP_ENV": "testing",
"BOOL_TEST": "true",
"INT_TEST": "42",
"DB_USER": "postgres",
"DB_PASS": "pass",
"DB_HOST": "localhost",
"DB_PORT": "5432",
"DB_NAME": "query_orchestration",
"DB_NOSSL": "true",
"AWS_ACCESS_KEY_ID": "key",
"AWS_SECRET_ACCESS_KEY": "secret",
"AWS_REGION": "region",
"SUB_FIELD1:": "value1",
"SUB_FIELD2": "42",
"PWD": "/foo",
},
wantErr: false,
},
{
name: "missing required env var",
envVars: map[string]string{
"APP_ENV": "testing",
"BOOL_TEST": "true",
"DB_USER": "postgres",
"DB_PASS": "pass",
"DB_HOST": "localhost",
"DB_PORT": "5432",
"DB_NAME": "query_orchestration",
"DB_NOSSL": "true",
"AWS_ACCESS_KEY_ID": "key",
"AWS_SECRET_ACCESS_KEY": "secret",
"AWS_REGION": "region",
"PWD": "/foo",
// INT_TEST intentionally omitted
},
wantErr: true,
errMessageContains: "INT_TEST",
},
{
name: "invalid boolean value",
envVars: map[string]string{
"APP_ENV": "testing",
"BOOL_TEST": "notabool",
"PWD": "/foo",
"INT_TEST": "42",
},
wantErr: true,
errMessageContains: "BoolTest",
},
{
name: "invalid integer value",
envVars: map[string]string{
"APP_ENV": "testing",
"BOOL_TEST": "true",
"PWD": "/foo",
"INT_TEST": "notanint",
},
wantErr: true,
errMessageContains: "parse error on field \"IntTest\"",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
os.Clearenv()
for k, v := range tt.envVars {
t.Setenv(k, v)
}
cfg := &TestingServiceNameConfig{}
err := InitializeConfig(cfg)
logger := cfg.GetLogger()
assert.NotNil(t, logger)
logger.Info("Logger initialized")
if tt.wantErr {
if err == nil {
t.Errorf("InitializeConfig() error = nil, wantErr %v", tt.wantErr)
return
}
// if we expect an error, check if the error message tt.errMessageContains is contained in the error message
if tt.errMessageContains != "" && !strings.Contains(err.Error(), tt.errMessageContains) {
t.Errorf("InitializeConfig() error = %v, want error containing %v", err, tt.errMessageContains)
}
} else {
if err != nil {
t.Errorf("InitializeConfig() unexpected error = %v", err)
return
}
// Verify the values were set correctly
if cfg.AppEnv != tt.envVars["APP_ENV"] {
t.Errorf("AppEnv = %v, want %v", cfg.AppEnv, tt.envVars["APP_ENV"])
}
if cfg.BoolTest != (tt.envVars["BOOL_TEST"] == "true") {
t.Errorf("BoolTest = %v, want %v", cfg.BoolTest, tt.envVars["BOOL_TEST"] == "true")
}
expectedInt := 42 // Known value from test cases
if cfg.IntTest != expectedInt {
t.Errorf("IntTest = %v, want %v", cfg.IntTest, expectedInt)
}
}
})
}
}
func TestGetBaseConfig(t *testing.T) {
assert.NotNil(t, getBaseConfig(&BaseConfig{}))
assert.Nil(t, getBaseConfig(nil))
assert.NotNil(t, getBaseConfig(&TestingServiceNameConfig{}))
}
func TestGetBasePath(t *testing.T) {
cfg := &BaseConfig{
Pwd: "pwd_path",
}
assert.Equal(t, "pwd_path", cfg.GetBasePath())
cfg.SetBasePath("base_path")
assert.Equal(t, "base_path", cfg.GetBasePath())
}
func TestSetBasePath(t *testing.T) {
cfg := &BaseConfig{}
cfg.SetBasePath("/i/am/here")
assert.Equal(t, "/i/am/here", cfg.BasePath)
}
func TestSetDBConfig(t *testing.T) {
cfg := &BaseConfig{}
dbcfg := database.DBConfig{
DBUser: "user",
}
cfg.SetDBConfig(&dbcfg)
assert.Equal(t, dbcfg, cfg.DBConfig)
}
func TestSetAWSConfig(t *testing.T) {
cfg := &BaseConfig{}
acfg := aws.AWSConfig{
AWSRegion: "user",
}
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)
}
}
}
}