Add per-subsystem config validation for wireguard, authelia, nftables, tunnel
- wireguard: ValidateConfig checks ListenPort range, Address/LanCIDR CIDR format. SaveConfig now validates before writing. - authelia: ValidateConfigOpts checks required fields (Domain, JWTSecret, SessionSecret) and StorageEncKey minimum length (20 chars). GenerateConfig now validates before generating. - nftables: ValidateWANInterface checks interface exists via net.InterfaceByName before generating rules that reference it. - tunnel: ValidateConfig checks TunnelID, PublicDomain, GatewayDomain are set and credentials file exists. WriteConfig now validates before generating.
This commit is contained in:
@@ -32,8 +32,28 @@ type ConfigOpts struct {
|
|||||||
SMTPPassword string // SMTP password (from secrets)
|
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 {
|
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 {
|
if err := m.EnsureDataDir(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package nftables
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -131,6 +132,18 @@ func (m *Manager) ValidateRules(rulesPath string) error {
|
|||||||
return nil
|
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.
|
// SafeApply validates, backs up, writes, and applies nftables rules.
|
||||||
// On apply failure, rolls back to the previous rules.
|
// On apply failure, rolls back to the previous rules.
|
||||||
func (m *Manager) SafeApply(content string) error {
|
func (m *Manager) SafeApply(content string) error {
|
||||||
|
|||||||
@@ -132,8 +132,32 @@ func (m *Manager) GenerateConfig(cfg Config, services []PublicDomain) string {
|
|||||||
return string(data)
|
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 {
|
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)
|
content := m.GenerateConfig(cfg, services)
|
||||||
if content == "" {
|
if content == "" {
|
||||||
// Remove config if tunnel is disabled or no public services
|
// Remove config if tunnel is disabled or no public services
|
||||||
|
|||||||
@@ -7,13 +7,17 @@ import (
|
|||||||
"testing"
|
"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{
|
return Config{
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
TunnelID: "abc-123",
|
TunnelID: "abc-123",
|
||||||
PublicDomain: "pub.payne.io",
|
PublicDomain: "pub.payne.io",
|
||||||
GatewayDomain: "payne.io",
|
GatewayDomain: "payne.io",
|
||||||
CredentialsDir: "/tmp/creds",
|
CredentialsDir: credsDir,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -25,7 +29,8 @@ func TestGenerateConfig_Basic(t *testing.T) {
|
|||||||
{Name: "dashboard"},
|
{Name: "dashboard"},
|
||||||
}
|
}
|
||||||
|
|
||||||
out := m.GenerateConfig(testConfig(), services)
|
cfg := testConfig(t)
|
||||||
|
out := m.GenerateConfig(cfg, services)
|
||||||
|
|
||||||
if out == "" {
|
if out == "" {
|
||||||
t.Fatal("expected non-empty config")
|
t.Fatal("expected non-empty config")
|
||||||
@@ -37,7 +42,7 @@ func TestGenerateConfig_Basic(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check credentials file
|
// 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)
|
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"},
|
{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") {
|
if !strings.Contains(out, "hostname: public-name.pub.payne.io") {
|
||||||
t.Errorf("expected subdomain override, got:\n%s", out)
|
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) {
|
func TestGenerateConfig_Disabled(t *testing.T) {
|
||||||
m := NewManager(t.TempDir())
|
m := NewManager(t.TempDir())
|
||||||
|
|
||||||
cfg := testConfig()
|
cfg := testConfig(t)
|
||||||
cfg.Enabled = false
|
cfg.Enabled = false
|
||||||
|
|
||||||
out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}})
|
out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}})
|
||||||
@@ -96,7 +101,7 @@ func TestGenerateConfig_Disabled(t *testing.T) {
|
|||||||
|
|
||||||
func TestGenerateConfig_NoServices(t *testing.T) {
|
func TestGenerateConfig_NoServices(t *testing.T) {
|
||||||
m := NewManager(t.TempDir())
|
m := NewManager(t.TempDir())
|
||||||
out := m.GenerateConfig(testConfig(), nil)
|
out := m.GenerateConfig(testConfig(t), nil)
|
||||||
if out != "" {
|
if out != "" {
|
||||||
t.Errorf("expected empty config with no services, got:\n%s", 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) {
|
func TestGenerateConfig_MissingTunnelID(t *testing.T) {
|
||||||
m := NewManager(t.TempDir())
|
m := NewManager(t.TempDir())
|
||||||
cfg := testConfig()
|
cfg := testConfig(t)
|
||||||
cfg.TunnelID = ""
|
cfg.TunnelID = ""
|
||||||
|
|
||||||
out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}})
|
out := m.GenerateConfig(cfg, []PublicDomain{{Name: "my-api"}})
|
||||||
@@ -118,7 +123,7 @@ func TestWriteConfig(t *testing.T) {
|
|||||||
m := NewManager(tmpDir)
|
m := NewManager(tmpDir)
|
||||||
|
|
||||||
services := []PublicDomain{{Name: "my-api"}}
|
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)
|
t.Fatalf("WriteConfig failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -139,12 +144,12 @@ func TestWriteConfig_RemovesWhenDisabled(t *testing.T) {
|
|||||||
m := NewManager(tmpDir)
|
m := NewManager(tmpDir)
|
||||||
|
|
||||||
// Write config first
|
// 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)
|
t.Fatalf("WriteConfig failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Now disable and write again — should remove the file
|
// Now disable and write again — should remove the file
|
||||||
cfg := testConfig()
|
cfg := testConfig(t)
|
||||||
cfg.Enabled = false
|
cfg.Enabled = false
|
||||||
if err := m.WriteConfig(cfg, []PublicDomain{{Name: "my-api"}}); err != nil {
|
if err := m.WriteConfig(cfg, []PublicDomain{{Name: "my-api"}}); err != nil {
|
||||||
t.Fatalf("WriteConfig (disable) failed: %v", err)
|
t.Fatalf("WriteConfig (disable) failed: %v", err)
|
||||||
|
|||||||
@@ -103,8 +103,29 @@ func (m *Manager) GetConfig() (*Config, error) {
|
|||||||
return &cfg, nil
|
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 {
|
func (m *Manager) SaveConfig(cfg *Config) error {
|
||||||
|
if err := ValidateConfig(cfg); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
if err := os.MkdirAll(m.vpnDir(), 0755); err != nil {
|
if err := os.MkdirAll(m.vpnDir(), 0755); err != nil {
|
||||||
return fmt.Errorf("create vpn dir: %w", err)
|
return fmt.Errorf("create vpn dir: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user