package config import ( "encoding/json" "fmt" "os" "path/filepath" "strings" "gopkg.in/yaml.v3" ) type Config struct { Settings Settings `yaml:"settings"` Vars map[string]string `yaml:"vars,omitempty"` PortGroups map[string]PortGroup `yaml:"portgroups"` Zones map[string]Zone `yaml:"zones"` Interfaces []Interface `yaml:"interfaces"` Hosts []Host `yaml:"hosts"` Policy []Policy `yaml:"policy"` Rules []Rule `yaml:"rules"` Blrules []BlruleRule `yaml:"blrules,omitempty"` SNAT []SNATRule `yaml:"snat"` StaticNAT []StaticNAT `yaml:"nat"` Netmap []Netmap `yaml:"netmap"` Providers []Provider `yaml:"providers"` Conntrack []ConntrackRule `yaml:"conntrack,omitempty"` Tunnels []Tunnel `yaml:"tunnels,omitempty"` RoutingRules []RoutingRule `yaml:"rtrules,omitempty"` StoppedRules []StoppedRule `yaml:"stoppedrules,omitempty"` ProxyARP []ProxyARP `yaml:"proxyarp,omitempty"` ProxyNDP []ProxyNDP `yaml:"proxyndp,omitempty"` Routes []StaticRoute `yaml:"routes,omitempty"` ArpRules []ArpRule `yaml:"arprules,omitempty"` Accounting []AccountingRule `yaml:"accounting,omitempty"` Mangle []MangleRule `yaml:"mangle,omitempty"` Maclist []MaclistEntry `yaml:"maclist,omitempty"` TCDevices []TCDevice `yaml:"tcdevices,omitempty"` TCClasses []TCClass `yaml:"tcclasses,omitempty"` TCFilters []TCFilter `yaml:"tcfilters,omitempty"` TCInterfaces []TCInterface `yaml:"tcinterfaces,omitempty"` TCPriorities []TCPriority `yaml:"tcpriority,omitempty"` Secmarks []SecmarkRule `yaml:"secmarks,omitempty"` } type AddressFamily string const ( FamilyINET AddressFamily = "inet" FamilyIP AddressFamily = "ip" FamilyIP6 AddressFamily = "ip6" ) type Settings struct { AddressFamily AddressFamily `yaml:"address_family,omitempty"` IPForwarding bool `yaml:"ip_forwarding"` LogLevel string `yaml:"log_level"` TableName string `yaml:"table_name"` // When true, auto-generate CONTINUE policies for sub-zones to their parent zones. ImplicitContinue bool `yaml:"implicit_continue,omitempty"` } // Load reads a config file in YAML or JSON format (detected by extension). func Load(path string) (*Config, error) { data, err := os.ReadFile(path) if err != nil { return nil, fmt.Errorf("reading config %s: %w", path, err) } var cfg Config ext := strings.ToLower(filepath.Ext(path)) switch ext { case ".json": if err := json.Unmarshal(data, &cfg); err != nil { return nil, fmt.Errorf("parsing JSON config: %w", err) } default: if err := yaml.Unmarshal(data, &cfg); err != nil { return nil, fmt.Errorf("parsing YAML config: %w", err) } } cfg.applyDefaults() return &cfg, nil } // ToYAML serializes the config to YAML bytes. func (c *Config) ToYAML() ([]byte, error) { return yaml.Marshal(c) } // ToJSON serializes the config to indented JSON bytes. func (c *Config) ToJSON() ([]byte, error) { return json.MarshalIndent(c, "", " ") } func (c *Config) applyDefaults() { if c.Settings.TableName == "" { c.Settings.TableName = "tomswall" } if c.Settings.LogLevel == "" { c.Settings.LogLevel = "info" } if c.Settings.AddressFamily == "" { c.Settings.AddressFamily = FamilyINET } } var validAddressFamilies = map[AddressFamily]bool{ FamilyINET: true, FamilyIP: true, FamilyIP6: true, } func (c *Config) validateSettings() error { if !validAddressFamilies[c.Settings.AddressFamily] { return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily) } return nil } func (c *Config) Validate() error { if err := c.validateSettings(); err != nil { return fmt.Errorf("settings: %w", err) } if err := c.validateZones(); err != nil { return fmt.Errorf("zones: %w", err) } if err := c.validateInterfaces(); err != nil { return fmt.Errorf("interfaces: %w", err) } if err := c.validateHosts(); err != nil { return fmt.Errorf("hosts: %w", err) } if err := c.validatePortGroups(); err != nil { return fmt.Errorf("portgroups: %w", err) } if err := c.validatePolicy(); err != nil { return fmt.Errorf("policy: %w", err) } if err := c.validateRules(); err != nil { return fmt.Errorf("rules: %w", err) } if err := c.validateSNAT(); err != nil { return fmt.Errorf("snat: %w", err) } if err := c.validateStaticNAT(); err != nil { return fmt.Errorf("nat: %w", err) } if err := c.validateNetmap(); err != nil { return fmt.Errorf("netmap: %w", err) } if err := c.validateProviders(); err != nil { return fmt.Errorf("providers: %w", err) } if err := c.validateVars(); err != nil { return fmt.Errorf("vars: %w", err) } if err := c.validateConntrack(); err != nil { return fmt.Errorf("conntrack: %w", err) } if err := c.validateBlrules(); err != nil { return fmt.Errorf("blrules: %w", err) } if err := c.validateTunnels(); err != nil { return fmt.Errorf("tunnels: %w", err) } if err := c.validateRoutingRules(); err != nil { return fmt.Errorf("rtrules: %w", err) } if err := c.validateStoppedRules(); err != nil { return fmt.Errorf("stoppedrules: %w", err) } if err := c.validateProxyARP(); err != nil { return fmt.Errorf("proxyarp: %w", err) } if err := c.validateProxyNDP(); err != nil { return fmt.Errorf("proxyndp: %w", err) } if err := c.validateRoutes(); err != nil { return fmt.Errorf("routes: %w", err) } if err := c.validateArpRules(); err != nil { return fmt.Errorf("arprules: %w", err) } if err := c.validateAccounting(); err != nil { return fmt.Errorf("accounting: %w", err) } if err := c.validateMangle(); err != nil { return fmt.Errorf("mangle: %w", err) } if err := c.validateMaclist(); err != nil { return fmt.Errorf("maclist: %w", err) } if err := c.validateTCDevices(); err != nil { return fmt.Errorf("tcdevices: %w", err) } if err := c.validateTCClasses(); err != nil { return fmt.Errorf("tcclasses: %w", err) } if err := c.validateTCFilters(); err != nil { return fmt.Errorf("tcfilters: %w", err) } if err := c.validateTCInterfaces(); err != nil { return fmt.Errorf("tcinterfaces: %w", err) } if err := c.validateTCPriority(); err != nil { return fmt.Errorf("tcpriority: %w", err) } if err := c.validateSecmarks(); err != nil { return fmt.Errorf("secmarks: %w", err) } return nil } func (c *Config) FirewallZone() string { for name, z := range c.Zones { if z.Type == ZoneFirewall { return name } } return "" } func (c *Config) ZoneInterfaces(zone string) []string { var ifaces []string for _, iface := range c.Interfaces { if iface.Zone == zone { ifaces = append(ifaces, iface.Interface) } } return ifaces } func (c *Config) ResolvePortGroup(name string) (*PortGroup, bool) { pg, ok := c.PortGroups[name] if !ok { return nil, false } return &pg, true }