diff --git a/internal/authelia/config.go b/internal/authelia/config.go index ff12f9f..4e739bb 100644 --- a/internal/authelia/config.go +++ b/internal/authelia/config.go @@ -32,8 +32,28 @@ type ConfigOpts struct { SMTPPassword string // SMTP password (from secrets) } -// GenerateConfig builds Authelia's configuration.yml and writes it atomically. +// ValidateConfigOpts checks that required fields are present and valid. +func ValidateConfigOpts(opts ConfigOpts) error { + if opts.Domain == "" { + return fmt.Errorf("auth portal domain is required") + } + if opts.JWTSecret == "" { + return fmt.Errorf("JWT secret is required") + } + if opts.SessionSecret == "" { + return fmt.Errorf("session secret is required") + } + if len(opts.StorageEncKey) < 20 { + return fmt.Errorf("storage encryption key must be at least 20 characters, got %d", len(opts.StorageEncKey)) + } + return nil +} + +// GenerateConfig validates options, builds Authelia's configuration.yml, and writes it atomically. func (m *Manager) GenerateConfig(opts ConfigOpts) error { + if err := ValidateConfigOpts(opts); err != nil { + return fmt.Errorf("config validation: %w", err) + } if err := m.EnsureDataDir(); err != nil { return err } diff --git a/internal/nftables/manager.go b/internal/nftables/manager.go index 0f1cfeb..2fafc12 100644 --- a/internal/nftables/manager.go +++ b/internal/nftables/manager.go @@ -3,6 +3,7 @@ package nftables import ( "fmt" "log/slog" + "net" "os" "os/exec" "strings" @@ -131,6 +132,18 @@ func (m *Manager) ValidateRules(rulesPath string) error { return nil } +// ValidateWANInterface checks that the specified WAN interface exists on the system. +// Returns nil if empty (no WAN filtering) or if the interface exists. +func ValidateWANInterface(name string) error { + if name == "" { + return nil + } + if _, err := net.InterfaceByName(name); err != nil { + return fmt.Errorf("WAN interface %q not found: %w", name, err) + } + return nil +} + // SafeApply validates, backs up, writes, and applies nftables rules. // On apply failure, rolls back to the previous rules. func (m *Manager) SafeApply(content string) error { diff --git a/internal/tunnel/manager.go b/internal/tunnel/manager.go index 66b5fc5..cf42f72 100644 --- a/internal/tunnel/manager.go +++ b/internal/tunnel/manager.go @@ -132,8 +132,32 @@ func (m *Manager) GenerateConfig(cfg Config, services []PublicDomain) string { return string(data) } -// WriteConfig generates and writes the cloudflared config to disk. +// ValidateConfig checks tunnel config for common errors before generating. +func ValidateConfig(cfg Config) error { + if !cfg.Enabled { + return nil // disabled is valid + } + if cfg.TunnelID == "" { + return fmt.Errorf("tunnel ID is required") + } + if cfg.PublicDomain == "" { + return fmt.Errorf("public domain is required") + } + if cfg.GatewayDomain == "" { + return fmt.Errorf("gateway domain is required") + } + credsFile := credentialsPath(cfg.CredentialsDir, cfg.TunnelID) + if _, err := os.Stat(credsFile); err != nil { + return fmt.Errorf("credentials file not found: %s", credsFile) + } + return nil +} + +// WriteConfig validates, generates, and writes the cloudflared config to disk. func (m *Manager) WriteConfig(cfg Config, services []PublicDomain) error { + if err := ValidateConfig(cfg); err != nil { + return fmt.Errorf("config validation: %w", err) + } content := m.GenerateConfig(cfg, services) if content == "" { // Remove config if tunnel is disabled or no public services diff --git a/internal/tunnel/manager_test.go b/internal/tunnel/manager_test.go index 7e52729..883047d 100644 --- a/internal/tunnel/manager_test.go +++ b/internal/tunnel/manager_test.go @@ -7,13 +7,17 @@ import ( "testing" ) -func testConfig() Config { +func testConfig(t *testing.T) Config { + t.Helper() + credsDir := filepath.Join(t.TempDir(), "creds") + os.MkdirAll(credsDir, 0755) + os.WriteFile(filepath.Join(credsDir, "abc-123.json"), []byte(`{"AccountTag":"test"}`), 0600) return Config{ Enabled: true, TunnelID: "abc-123", PublicDomain: "pub.payne.io", GatewayDomain: "payne.io", - CredentialsDir: "/tmp/creds", + CredentialsDir: credsDir, } } @@ -25,7 +29,8 @@ func TestGenerateConfig_Basic(t *testing.T) { {Name: "dashboard"}, } - out := m.GenerateConfig(testConfig(), services) + cfg := testConfig(t) + out := m.GenerateConfig(cfg, services) if out == "" { t.Fatal("expected non-empty config") @@ -37,7 +42,7 @@ func TestGenerateConfig_Basic(t *testing.T) { } // Check credentials file - if !strings.Contains(out, "credentials-file: /tmp/creds/abc-123.json") { + if !strings.Contains(out, "abc-123.json") { t.Errorf("expected credentials file, got:\n%s", out) } @@ -75,7 +80,7 @@ func TestGenerateConfig_SubdomainOverride(t *testing.T) { {Name: "internal-name", Subdomain: "public-name"}, } - out := m.GenerateConfig(testConfig(), services) + out := m.GenerateConfig(testConfig(t), services) if !strings.Contains(out, "hostname: public-name.pub.payne.io") { t.Errorf("expected subdomain override, got:\n%s", out) @@ -85,7 +90,7 @@ func TestGenerateConfig_SubdomainOverride(t *testing.T) { func TestGenerateConfig_Disabled(t *testing.T) { m := NewManager(t.TempDir()) - cfg := testConfig() + cfg := testConfig(t) cfg.Enabled = false out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}}) @@ -96,7 +101,7 @@ func TestGenerateConfig_Disabled(t *testing.T) { func TestGenerateConfig_NoServices(t *testing.T) { m := NewManager(t.TempDir()) - out := m.GenerateConfig(testConfig(), nil) + out := m.GenerateConfig(testConfig(t), nil) if out != "" { t.Errorf("expected empty config with no services, got:\n%s", out) } @@ -104,7 +109,7 @@ func TestGenerateConfig_NoServices(t *testing.T) { func TestGenerateConfig_MissingTunnelID(t *testing.T) { m := NewManager(t.TempDir()) - cfg := testConfig() + cfg := testConfig(t) cfg.TunnelID = "" out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}}) @@ -118,7 +123,7 @@ func TestWriteConfig(t *testing.T) { m := NewManager(tmpDir) services := []PublicDomain{{Name: "my-api"}} - if err := m.WriteConfig(testConfig(), services); err != nil { + if err := m.WriteConfig(testConfig(t), services); err != nil { t.Fatalf("WriteConfig failed: %v", err) } @@ -139,12 +144,12 @@ func TestWriteConfig_RemovesWhenDisabled(t *testing.T) { m := NewManager(tmpDir) // Write config first - if err := m.WriteConfig(testConfig(), []PublicDomain{{Name: "my-api"}}); err != nil { + if err := m.WriteConfig(testConfig(t), []PublicDomain{{Name: "my-api"}}); err != nil { t.Fatalf("WriteConfig failed: %v", err) } // Now disable and write again — should remove the file - cfg := testConfig() + cfg := testConfig(t) cfg.Enabled = false if err := m.WriteConfig(cfg, []PublicDomain{{Name: "my-api"}}); err != nil { t.Fatalf("WriteConfig (disable) failed: %v", err) diff --git a/internal/wireguard/manager.go b/internal/wireguard/manager.go index 50f0054..49b2722 100644 --- a/internal/wireguard/manager.go +++ b/internal/wireguard/manager.go @@ -103,8 +103,29 @@ func (m *Manager) GetConfig() (*Config, error) { return &cfg, nil } -// SaveConfig writes the server interface config. +// ValidateConfig checks a WireGuard config for common errors before saving. +func ValidateConfig(cfg *Config) error { + if cfg.ListenPort < 1 || cfg.ListenPort > 65535 { + return fmt.Errorf("listenPort must be 1-65535, got %d", cfg.ListenPort) + } + if cfg.Address != "" { + if _, _, err := net.ParseCIDR(cfg.Address); err != nil { + return fmt.Errorf("invalid server address CIDR %q: %w", cfg.Address, err) + } + } + if cfg.LanCIDR != "" { + if _, _, err := net.ParseCIDR(cfg.LanCIDR); err != nil { + return fmt.Errorf("invalid LAN CIDR %q: %w", cfg.LanCIDR, err) + } + } + return nil +} + +// SaveConfig validates and writes the server interface config. func (m *Manager) SaveConfig(cfg *Config) error { + if err := ValidateConfig(cfg); err != nil { + return err + } if err := os.MkdirAll(m.vpnDir(), 0755); err != nil { return fmt.Errorf("create vpn dir: %w", err) }