You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

484 lines
11 KiB
Go

package utils
import (
"os"
"testing"
"time"
)
func TestGetDatabaseConfig(t *testing.T) {
// Clean up environment before test
cleanupDBEnv := func() {
os.Unsetenv("DATABASE_URL")
os.Unsetenv("DB_HOST")
os.Unsetenv("DB_PORT")
os.Unsetenv("DB_USER")
os.Unsetenv("DB_PASSWORD")
os.Unsetenv("DB_NAME")
os.Unsetenv("DB_SSLMODE")
}
t.Run("Default configuration", func(t *testing.T) {
cleanupDBEnv()
config := GetDatabaseConfig()
if config.URL != "sqlite://./app.db" {
t.Errorf("Expected default URL 'sqlite://./app.db', got '%s'", config.URL)
}
if config.Driver != "sqlite3" {
t.Errorf("Expected driver 'sqlite3', got '%s'", config.Driver)
}
if config.Host != "localhost" {
t.Errorf("Expected host 'localhost', got '%s'", config.Host)
}
if config.Port != 5432 {
t.Errorf("Expected port 5432, got %d", config.Port)
}
})
t.Run("PostgreSQL configuration", func(t *testing.T) {
cleanupDBEnv()
os.Setenv("DATABASE_URL", "postgres://user:pass@host:5433/dbname")
os.Setenv("DB_HOST", "custom-host")
os.Setenv("DB_PORT", "5433")
os.Setenv("DB_USER", "custom-user")
os.Setenv("DB_PASSWORD", "custom-pass")
os.Setenv("DB_NAME", "custom-db")
os.Setenv("DB_SSLMODE", "require")
defer cleanupDBEnv()
config := GetDatabaseConfig()
if config.Driver != "postgres" {
t.Errorf("Expected driver 'postgres', got '%s'", config.Driver)
}
if config.Host != "custom-host" {
t.Errorf("Expected host 'custom-host', got '%s'", config.Host)
}
if config.Port != 5433 {
t.Errorf("Expected port 5433, got %d", config.Port)
}
if config.User != "custom-user" {
t.Errorf("Expected user 'custom-user', got '%s'", config.User)
}
if config.SSLMode != "require" {
t.Errorf("Expected sslmode 'require', got '%s'", config.SSLMode)
}
})
t.Run("Unknown database URL defaults to sqlite3", func(t *testing.T) {
cleanupDBEnv()
os.Setenv("DATABASE_URL", "unknown://user:pass@host:1234/dbname")
defer cleanupDBEnv()
config := GetDatabaseConfig()
if config.Driver != "sqlite3" {
t.Errorf("Expected driver 'sqlite3' for unknown URL, got '%s'", config.Driver)
}
if config.URL != "unknown://user:pass@host:1234/dbname" {
t.Errorf("URL should be preserved as-is, got '%s'", config.URL)
}
})
}
func TestGetServerConfig(t *testing.T) {
cleanupServerEnv := func() {
os.Unsetenv("PORT")
os.Unsetenv("HOST")
os.Unsetenv("READ_TIMEOUT")
os.Unsetenv("WRITE_TIMEOUT")
os.Unsetenv("SHUTDOWN_TIMEOUT")
os.Unsetenv("LOG_LEVEL")
}
t.Run("Default configuration", func(t *testing.T) {
cleanupServerEnv()
config := GetServerConfig("8080")
if config.Port != "8080" {
t.Errorf("Expected port '8080', got '%s'", config.Port)
}
if config.Host != "" {
t.Errorf("Expected empty host, got '%s'", config.Host)
}
if config.ReadTimeout != 30*time.Second {
t.Errorf("Expected read timeout 30s, got %v", config.ReadTimeout)
}
if config.LogLevel != "info" {
t.Errorf("Expected log level 'info', got '%s'", config.LogLevel)
}
})
t.Run("Custom configuration", func(t *testing.T) {
cleanupServerEnv()
os.Setenv("PORT", "9000")
os.Setenv("HOST", "127.0.0.1")
os.Setenv("READ_TIMEOUT", "45s")
os.Setenv("WRITE_TIMEOUT", "60s")
os.Setenv("SHUTDOWN_TIMEOUT", "10s")
os.Setenv("LOG_LEVEL", "debug")
defer cleanupServerEnv()
config := GetServerConfig("8080")
if config.Port != "9000" {
t.Errorf("Expected port '9000', got '%s'", config.Port)
}
if config.Host != "127.0.0.1" {
t.Errorf("Expected host '127.0.0.1', got '%s'", config.Host)
}
if config.ReadTimeout != 45*time.Second {
t.Errorf("Expected read timeout 45s, got %v", config.ReadTimeout)
}
if config.WriteTimeout != 60*time.Second {
t.Errorf("Expected write timeout 60s, got %v", config.WriteTimeout)
}
if config.ShutdownTimeout != 10*time.Second {
t.Errorf("Expected shutdown timeout 10s, got %v", config.ShutdownTimeout)
}
if config.LogLevel != "debug" {
t.Errorf("Expected log level 'debug', got '%s'", config.LogLevel)
}
})
}
func TestServerConfigGetListenAddress(t *testing.T) {
tests := []struct {
name string
host string
port string
expected string
}{
{
name: "With host",
host: "127.0.0.1",
port: "8080",
expected: "127.0.0.1:8080",
},
{
name: "Without host",
host: "",
port: "8080",
expected: ":8080",
},
{
name: "With localhost",
host: "localhost",
port: "3000",
expected: "localhost:3000",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
config := ServerConfig{
Host: tt.host,
Port: tt.port,
}
result := config.GetListenAddress()
if result != tt.expected {
t.Errorf("GetListenAddress() = %v, want %v", result, tt.expected)
}
})
}
}
func TestGetAuthConfig(t *testing.T) {
cleanupAuthEnv := func() {
os.Unsetenv("JWT_SECRET")
os.Unsetenv("TOKEN_EXPIRATION")
os.Unsetenv("JWT_ISSUER")
os.Unsetenv("REQUIRE_MFA")
}
t.Run("Default configuration", func(t *testing.T) {
cleanupAuthEnv()
config := GetAuthConfig("test-service")
if config.JWTSecret != "dev-secret-test-service" {
t.Errorf("Expected JWT secret 'dev-secret-test-service', got '%s'", config.JWTSecret)
}
if config.TokenExpiration != 24*time.Hour {
t.Errorf("Expected token expiration 24h, got %v", config.TokenExpiration)
}
if config.Issuer != "knowfoolery-test-service" {
t.Errorf("Expected issuer 'knowfoolery-test-service', got '%s'", config.Issuer)
}
if config.RequireMFA != false {
t.Errorf("Expected require MFA false, got %v", config.RequireMFA)
}
})
t.Run("Custom configuration", func(t *testing.T) {
cleanupAuthEnv()
os.Setenv("JWT_SECRET", "custom-secret")
os.Setenv("TOKEN_EXPIRATION", "2h")
os.Setenv("JWT_ISSUER", "custom-issuer")
os.Setenv("REQUIRE_MFA", "true")
defer cleanupAuthEnv()
config := GetAuthConfig("test-service")
if config.JWTSecret != "custom-secret" {
t.Errorf("Expected JWT secret 'custom-secret', got '%s'", config.JWTSecret)
}
if config.TokenExpiration != 2*time.Hour {
t.Errorf("Expected token expiration 2h, got %v", config.TokenExpiration)
}
if config.Issuer != "custom-issuer" {
t.Errorf("Expected issuer 'custom-issuer', got '%s'", config.Issuer)
}
if config.RequireMFA != true {
t.Errorf("Expected require MFA true, got %v", config.RequireMFA)
}
})
}
func TestEnvironmentChecks(t *testing.T) {
cleanupEnv := func() {
os.Unsetenv("ENVIRONMENT")
}
tests := []struct {
name string
envValue string
isProd bool
isDev bool
getEnv string
}{
{
name: "Production environment",
envValue: "production",
isProd: true,
isDev: false,
getEnv: "production",
},
{
name: "Prod environment",
envValue: "prod",
isProd: true,
isDev: false,
getEnv: "prod",
},
{
name: "Development environment",
envValue: "development",
isProd: false,
isDev: true,
getEnv: "development",
},
{
name: "Dev environment",
envValue: "dev",
isProd: false,
isDev: true,
getEnv: "dev",
},
{
name: "Default environment",
envValue: "",
isProd: false,
isDev: true,
getEnv: "development",
},
{
name: "Custom environment",
envValue: "staging",
isProd: false,
isDev: false,
getEnv: "staging",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cleanupEnv()
if tt.envValue != "" {
os.Setenv("ENVIRONMENT", tt.envValue)
defer cleanupEnv()
}
if IsProduction() != tt.isProd {
t.Errorf("IsProduction() = %v, want %v", IsProduction(), tt.isProd)
}
if IsDevelopment() != tt.isDev {
t.Errorf("IsDevelopment() = %v, want %v", IsDevelopment(), tt.isDev)
}
if GetEnvironment() != tt.getEnv {
t.Errorf("GetEnvironment() = %v, want %v", GetEnvironment(), tt.getEnv)
}
})
}
}
func TestGetServiceName(t *testing.T) {
cleanupEnv := func() {
os.Unsetenv("SERVICE_NAME")
}
t.Run("Default service name", func(t *testing.T) {
cleanupEnv()
result := GetServiceName("default-service")
if result != "default-service" {
t.Errorf("Expected 'default-service', got '%s'", result)
}
})
t.Run("Custom service name", func(t *testing.T) {
cleanupEnv()
os.Setenv("SERVICE_NAME", "custom-service")
defer cleanupEnv()
result := GetServiceName("default-service")
if result != "custom-service" {
t.Errorf("Expected 'custom-service', got '%s'", result)
}
})
}
func TestGetVersion(t *testing.T) {
cleanupEnv := func() {
os.Unsetenv("VERSION")
}
t.Run("Default version", func(t *testing.T) {
cleanupEnv()
result := GetVersion()
if result != "development" {
t.Errorf("Expected 'development', got '%s'", result)
}
})
t.Run("Custom version", func(t *testing.T) {
cleanupEnv()
os.Setenv("VERSION", "1.2.3")
defer cleanupEnv()
result := GetVersion()
if result != "1.2.3" {
t.Errorf("Expected '1.2.3', got '%s'", result)
}
})
}
func TestLogSafeMethods(t *testing.T) {
t.Run("DatabaseConfig LogSafe", func(t *testing.T) {
config := DatabaseConfig{
Driver: "postgres",
Host: "localhost",
Port: 5432,
User: "testuser",
Password: "secret-password",
Database: "testdb",
SSLMode: "disable",
}
safe := config.LogSafe()
// Check that all non-sensitive fields are present
if safe["driver"] != "postgres" {
t.Error("Driver not present in safe log")
}
if safe["host"] != "localhost" {
t.Error("Host not present in safe log")
}
if safe["port"] != 5432 {
t.Error("Port not present in safe log")
}
// Check that password is NOT present
if _, exists := safe["password"]; exists {
t.Error("Password should not be present in safe log")
}
})
t.Run("ServerConfig LogSafe", func(t *testing.T) {
config := ServerConfig{
Port: "8080",
Host: "localhost",
ReadTimeout: 30 * time.Second,
WriteTimeout: 45 * time.Second,
LogLevel: "debug",
}
safe := config.LogSafe()
if safe["port"] != "8080" {
t.Error("Port not present in safe log")
}
if safe["log_level"] != "debug" {
t.Error("Log level not present in safe log")
}
})
t.Run("AuthConfig LogSafe", func(t *testing.T) {
config := AuthConfig{
JWTSecret: "secret-jwt-key",
TokenExpiration: 24 * time.Hour,
Issuer: "test-issuer",
RequireMFA: true,
}
safe := config.LogSafe()
// Check that non-sensitive fields are present
if safe["issuer"] != "test-issuer" {
t.Error("Issuer not present in safe log")
}
if safe["require_mfa"] != true {
t.Error("RequireMFA not present in safe log")
}
// Check that JWT secret is NOT present
if _, exists := safe["jwt_secret"]; exists {
t.Error("JWT secret should not be present in safe log")
}
})
}