Skip to content
Open
Show file tree
Hide file tree
Changes from 4 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 71 additions & 0 deletions config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,12 @@ package config
import (
"errors"
"flag"
"fmt"
"net"
"net/url"
"os"
"path/filepath"
"strings"
"time"

"github.com/ilyakaznacheev/cleanenv"
Expand All @@ -17,6 +20,14 @@ var ConsoleConfig *Config
// TrayMode indicates whether to run with system tray UI.
var TrayMode bool

// Validation errors for configuration.
var (
ErrSecretsAddrInsecure = errors.New("SECRETS_ADDR must use HTTPS for non-localhost addresses")
ErrSecretsAddrInvalid = errors.New("invalid SECRETS_ADDR")
ErrSecretsAddrMissingScheme = errors.New("SECRETS_ADDR missing scheme (use http:// or https://)")
ErrSecretsAddrNoHost = errors.New("SECRETS_ADDR contains no host")
)

const defaultHost = "localhost"

// DefaultSessionCookieName names the HttpOnly cookie holding the session JWT.
Expand Down Expand Up @@ -365,5 +376,65 @@ func NewConfig() (*Config, error) {
return nil, err
}

if err := ConsoleConfig.Validate(); err != nil {
return nil, err
}

return ConsoleConfig, nil
}

// Validate checks configuration for security and correctness issues.
func (c *Config) Validate() error {
// Ensure non-localhost Vault addresses use HTTPS
if c.Secrets.Address != "" { //nolint:staticcheck // QF1008: explicit field reference is clearer
Comment thread
nmgaston marked this conversation as resolved.
if err := c.validateSecretsAddr(); err != nil {
return err
}
}

return nil
}

// validateSecretsAddr ensures Vault address uses HTTPS (except for localhost).
func (c *Config) validateSecretsAddr() error {
parsed, err := url.Parse(c.Secrets.Address) //nolint:staticcheck // QF1008: explicit field reference is clearer
if err != nil {
return fmt.Errorf("%w: %w", ErrSecretsAddrInvalid, err)
}

// Check for valid scheme (must be http or https)
if parsed.Scheme != "http" && parsed.Scheme != "https" {
if !strings.Contains(c.Secrets.Address, "://") { //nolint:staticcheck // QF1008: explicit field reference is clearer
return fmt.Errorf("%w: %q", ErrSecretsAddrMissingScheme, c.Secrets.Address) //nolint:staticcheck // QF1008: explicit field reference is clearer
}

return fmt.Errorf("%w: unsupported scheme %q in %q", ErrSecretsAddrInvalid, parsed.Scheme, c.Secrets.Address) //nolint:staticcheck // QF1008: explicit field reference is clearer
}
Comment thread
nmgaston marked this conversation as resolved.

// Check host is present before checking HTTPS requirement
hostname := parsed.Hostname()
if hostname == "" {
return fmt.Errorf("%w: %q", ErrSecretsAddrNoHost, c.Secrets.Address) //nolint:staticcheck // QF1008: explicit field reference is clearer
}

// Enforce HTTPS for non-localhost
if parsed.Scheme == "http" && !isLocalhost(hostname) {
return ErrSecretsAddrInsecure
}

return nil
}

func isLocalhost(host string) bool {
host = strings.TrimSuffix(host, ".")

if strings.EqualFold(host, defaultHost) {
return true
}

if ip := net.ParseIP(host); ip != nil {
return ip.IsLoopback()
}

return false
}
Comment thread
nmgaston marked this conversation as resolved.
234 changes: 234 additions & 0 deletions config/config_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package config

import (
"errors"
"os"
"testing"

Expand All @@ -13,6 +14,19 @@ func clearEnv() {
os.Unsetenv("LOG_LEVEL")
os.Unsetenv("DB_POOL_MAX")
os.Unsetenv("DB_URL")
os.Unsetenv("SECRETS_ADDR")
}

func TestNewConfig_InvalidEnvVar(t *testing.T) {
clearEnv()
defer clearEnv()

// DB_POOL_MAX expects an int; a non-numeric value causes cleanenv.ReadEnv to fail.
t.Setenv("DB_POOL_MAX", "not-a-number")

cfg, err := NewConfig()
assert.Error(t, err)
assert.Nil(t, cfg)
}

func TestNewConfig_Defaults(t *testing.T) { //nolint:paralleltest // cannot have simultaneous tests modifying environment variables
Expand Down Expand Up @@ -107,3 +121,223 @@ postgres:
assert.Equal(t, 10, cfg.PoolMax)
assert.Equal(t, "postgres://envuser:envpassword@localhost:5432/envdb", cfg.DB.URL)
}

func TestValidate_SecretsAddrEmpty(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: "",
},
}
assert.NoError(t, cfg.Validate())
}

func TestValidate_SecretsAddrHTTPSNonLocalhost(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: "https://vault.example.com:8200",
},
}
assert.NoError(t, cfg.Validate())
}

func TestValidate_SecretsAddrHTTPLocalhost(t *testing.T) {
t.Parallel()

testCases := []string{
"http://localhost:8200",
"http://127.0.0.1:8200",
"http://127.0.0.2:8200",
"http://[::1]:8200",
"http://[::1]",
}
for _, addr := range testCases {
t.Run(addr, func(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: addr,
},
}
assert.NoError(t, cfg.Validate(), "expected %s to be valid", addr)
})
}
}

func TestValidate_SecretsAddrHTTPNonLocalhost(t *testing.T) {
t.Parallel()

testCases := []string{
"http://vault.example.com:8200",
"http://192.168.1.1:8200",
"http://vault-server:8200",
}
for _, addr := range testCases {
t.Run(addr, func(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: addr,
},
}
err := cfg.Validate()
assert.Error(t, err, "expected %s to fail", addr)
assert.True(t, errors.Is(err, ErrSecretsAddrInsecure), "expected ErrSecretsAddrInsecure, got %v", err)
})
}
}

func TestValidate_SecretsAddrInvalidURL(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: "://invalid",
},
}
err := cfg.Validate()
assert.Error(t, err)
assert.True(t, errors.Is(err, ErrSecretsAddrInvalid), "expected ErrSecretsAddrInvalid, got %v", err)
}

func TestValidate_SecretsAddrMissingScheme(t *testing.T) {
t.Parallel()

testCases := []string{
"vault.example.com:8200",
"localhost:8200",
}

for _, addr := range testCases {
t.Run(addr, func(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: addr,
},
}

err := cfg.Validate()
assert.Error(t, err)
assert.True(t, errors.Is(err, ErrSecretsAddrMissingScheme), "expected ErrSecretsAddrMissingScheme, got %v", err)
})
}
}

func TestValidate_SecretsAddrUnsupportedScheme(t *testing.T) {
t.Parallel()

testCases := []string{
"ftp://vault.example.com:8200",
"file:///tmp/vault",
}

for _, addr := range testCases {
t.Run(addr, func(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: addr,
},
}

err := cfg.Validate()
assert.Error(t, err)
assert.True(t, errors.Is(err, ErrSecretsAddrInvalid), "expected ErrSecretsAddrInvalid, got %v", err)
assert.Contains(t, err.Error(), "unsupported scheme")
})
}
}

func TestIsLocalhost(t *testing.T) {
t.Parallel()

testCases := []struct {
host string
expected bool
}{
// Localhost variants
{"localhost", true},
{"LOCALHOST", true},
{"localhost.", true},
{"LOCALHOST.", true},
{"127.0.0.1", true},
{"127.255.255.255", true},
{"::1", true},
{"127.1.1.1", true},

// Non-localhost
{"192.168.1.1", false},
{"vault.example.com", false},
{"172.16.0.1", false},
{"example.com", false},
{"2001:db8::1", false},
}
for _, tc := range testCases {
t.Run(tc.host, func(t *testing.T) {
t.Parallel()

result := isLocalhost(tc.host)
assert.Equal(t, tc.expected, result, "isLocalhost(%s) = %v, want %v", tc.host, result, tc.expected)
})
}
}

func TestValidate_CallsValidateSecretsAddr(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: "http://vault.example.com:8200", // Remote HTTP - should fail
},
}

err := cfg.Validate()
assert.Error(t, err)
assert.True(t, errors.Is(err, ErrSecretsAddrInsecure), "expected ErrSecretsAddrInsecure, got %v", err)
}

func TestValidate_AllowsValidRemoteHTTPS(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: "https://vault.example.com:8200",
},
}

err := cfg.Validate()
assert.NoError(t, err)
}

func TestValidateSecretsAddr_NoHost(t *testing.T) {
t.Parallel()

testCases := []string{
"http://", // Scheme but no host
"https://", // Scheme but no host
"http://:8200", // Scheme with port but no host
"https://:8200", // Scheme with port but no host
}
for _, addr := range testCases {
t.Run(addr, func(t *testing.T) {
t.Parallel()

cfg := &Config{
Secrets: Secrets{
Address: addr,
},
}
err := cfg.validateSecretsAddr()
assert.Error(t, err, "expected %s to fail", addr)
assert.True(t, errors.Is(err, ErrSecretsAddrNoHost), "expected ErrSecretsAddrNoHost, got %v", err)
})
}
}
Loading