Initial scaffold for tomswall
Spiritual successor to shorewall — manages nftables directly via google/nftables. Reads a single YAML config covering zones, interfaces, hosts, policy, rules, snat, and named portgroups. Computes differential changes against the running nftables state and applies them atomically. Supports detecting and purging rules added outside of tomswall.
This commit is contained in:
@@ -0,0 +1,101 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Settings Settings `yaml:"settings"`
|
||||
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"`
|
||||
SNAT []SNATRule `yaml:"snat"`
|
||||
}
|
||||
|
||||
type Settings struct {
|
||||
IPForwarding bool `yaml:"ip_forwarding"`
|
||||
LogLevel string `yaml:"log_level"`
|
||||
TableName string `yaml:"table_name"`
|
||||
}
|
||||
|
||||
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
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("parsing config: %w", err)
|
||||
}
|
||||
|
||||
cfg.applyDefaults()
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
func (c *Config) applyDefaults() {
|
||||
if c.Settings.TableName == "" {
|
||||
c.Settings.TableName = "tomswall"
|
||||
}
|
||||
if c.Settings.LogLevel == "" {
|
||||
c.Settings.LogLevel = "info"
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
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)
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type Host struct {
|
||||
Zone string `yaml:"zone"`
|
||||
Interface string `yaml:"interface"`
|
||||
Addresses []string `yaml:"addresses"`
|
||||
Options []string `yaml:"options,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Config) validateHosts() error {
|
||||
for i, h := range c.Hosts {
|
||||
if h.Zone == "" {
|
||||
return fmt.Errorf("host[%d]: zone required", i)
|
||||
}
|
||||
if _, ok := c.Zones[h.Zone]; !ok {
|
||||
return fmt.Errorf("host[%d]: zone %q not defined", i, h.Zone)
|
||||
}
|
||||
if h.Interface == "" {
|
||||
return fmt.Errorf("host[%d]: interface required", i)
|
||||
}
|
||||
if len(h.Addresses) == 0 {
|
||||
return fmt.Errorf("host[%d]: at least one address required", i)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type Interface struct {
|
||||
Zone string `yaml:"zone"`
|
||||
Interface string `yaml:"interface"`
|
||||
Options InterfaceOptions `yaml:"options,omitempty"`
|
||||
}
|
||||
|
||||
type InterfaceOptions struct {
|
||||
DHCP bool `yaml:"dhcp,omitempty"`
|
||||
TCPFlags bool `yaml:"tcpflags,omitempty"`
|
||||
NoSmurfs bool `yaml:"nosmurfs,omitempty"`
|
||||
RouteBack bool `yaml:"routeback,omitempty"`
|
||||
Bridge bool `yaml:"bridge,omitempty"`
|
||||
Optional bool `yaml:"optional,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Config) validateInterfaces() error {
|
||||
seen := make(map[string]bool)
|
||||
for i, iface := range c.Interfaces {
|
||||
if iface.Interface == "" {
|
||||
return fmt.Errorf("interface[%d]: interface name required", i)
|
||||
}
|
||||
if iface.Zone == "" {
|
||||
return fmt.Errorf("interface[%d] %q: zone required", i, iface.Interface)
|
||||
}
|
||||
if _, ok := c.Zones[iface.Zone]; !ok {
|
||||
return fmt.Errorf("interface[%d] %q: zone %q not defined", i, iface.Interface, iface.Zone)
|
||||
}
|
||||
if seen[iface.Interface] {
|
||||
return fmt.Errorf("interface[%d]: duplicate interface %q", i, iface.Interface)
|
||||
}
|
||||
seen[iface.Interface] = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type PolicyAction string
|
||||
|
||||
const (
|
||||
PolicyAccept PolicyAction = "accept"
|
||||
PolicyDrop PolicyAction = "drop"
|
||||
PolicyReject PolicyAction = "reject"
|
||||
PolicyContinue PolicyAction = "continue"
|
||||
PolicyNone PolicyAction = "none"
|
||||
)
|
||||
|
||||
type Policy struct {
|
||||
Source string `yaml:"source"`
|
||||
Dest string `yaml:"dest"`
|
||||
Action PolicyAction `yaml:"action"`
|
||||
Log string `yaml:"log,omitempty"`
|
||||
RateLimit string `yaml:"rate_limit,omitempty"`
|
||||
ConnLimit int `yaml:"conn_limit,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Config) validatePolicy() error {
|
||||
if len(c.Policy) == 0 {
|
||||
return fmt.Errorf("no policies defined")
|
||||
}
|
||||
|
||||
for i, p := range c.Policy {
|
||||
if p.Source == "" {
|
||||
return fmt.Errorf("policy[%d]: source required", i)
|
||||
}
|
||||
if p.Dest == "" {
|
||||
return fmt.Errorf("policy[%d]: dest required", i)
|
||||
}
|
||||
if p.Source != "all" {
|
||||
if _, ok := c.Zones[p.Source]; !ok {
|
||||
return fmt.Errorf("policy[%d]: source zone %q not defined", i, p.Source)
|
||||
}
|
||||
}
|
||||
if p.Dest != "all" {
|
||||
if _, ok := c.Zones[p.Dest]; !ok {
|
||||
return fmt.Errorf("policy[%d]: dest zone %q not defined", i, p.Dest)
|
||||
}
|
||||
}
|
||||
switch p.Action {
|
||||
case PolicyAccept, PolicyDrop, PolicyReject, PolicyContinue, PolicyNone:
|
||||
default:
|
||||
return fmt.Errorf("policy[%d]: unknown action %q", i, p.Action)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type PortGroup struct {
|
||||
Proto string `yaml:"proto"`
|
||||
Ports PortSpec `yaml:"ports"`
|
||||
}
|
||||
|
||||
// ParsedPorts returns individual port numbers and ranges as (start, end) pairs.
|
||||
func (pg *PortGroup) ParsedPorts() (singles []uint16, ranges [][2]uint16, err error) {
|
||||
for _, p := range pg.Ports {
|
||||
if strings.Contains(p, "-") {
|
||||
parts := strings.SplitN(p, "-", 2)
|
||||
start, err := strconv.ParseUint(parts[0], 10, 16)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("invalid port range start %q: %w", parts[0], err)
|
||||
}
|
||||
end, err := strconv.ParseUint(parts[1], 10, 16)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("invalid port range end %q: %w", parts[1], err)
|
||||
}
|
||||
if start > end {
|
||||
return nil, nil, fmt.Errorf("port range %d-%d: start > end", start, end)
|
||||
}
|
||||
ranges = append(ranges, [2]uint16{uint16(start), uint16(end)})
|
||||
} else {
|
||||
port, err := strconv.ParseUint(p, 10, 16)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("invalid port %q: %w", p, err)
|
||||
}
|
||||
singles = append(singles, uint16(port))
|
||||
}
|
||||
}
|
||||
return singles, ranges, nil
|
||||
}
|
||||
|
||||
func (c *Config) validatePortGroups() error {
|
||||
for name, pg := range c.PortGroups {
|
||||
if pg.Proto == "" {
|
||||
return fmt.Errorf("portgroup %q: proto required", name)
|
||||
}
|
||||
if pg.Proto != "tcp" && pg.Proto != "udp" {
|
||||
return fmt.Errorf("portgroup %q: proto must be tcp or udp, got %q", name, pg.Proto)
|
||||
}
|
||||
if len(pg.Ports) == 0 {
|
||||
return fmt.Errorf("portgroup %q: at least one port required", name)
|
||||
}
|
||||
if _, _, err := pg.ParsedPorts(); err != nil {
|
||||
return fmt.Errorf("portgroup %q: %w", name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type RuleAction string
|
||||
|
||||
const (
|
||||
RuleAccept RuleAction = "accept"
|
||||
RuleDrop RuleAction = "drop"
|
||||
RuleReject RuleAction = "reject"
|
||||
RuleDNAT RuleAction = "dnat"
|
||||
RuleRedirect RuleAction = "redirect"
|
||||
RuleLog RuleAction = "log"
|
||||
)
|
||||
|
||||
type Rule struct {
|
||||
Action RuleAction `yaml:"action"`
|
||||
Source string `yaml:"source"`
|
||||
Dest string `yaml:"dest"`
|
||||
Proto string `yaml:"proto,omitempty"`
|
||||
DPort PortSpec `yaml:"dport,omitempty"`
|
||||
SPort PortSpec `yaml:"sport,omitempty"`
|
||||
PortGroup string `yaml:"portgroup,omitempty"`
|
||||
Log string `yaml:"log,omitempty"`
|
||||
DNATDest string `yaml:"dnat_dest,omitempty"`
|
||||
RateLimit string `yaml:"rate_limit,omitempty"`
|
||||
ConnLimit int `yaml:"conn_limit,omitempty"`
|
||||
Comment string `yaml:"comment,omitempty"`
|
||||
}
|
||||
|
||||
// PortSpec supports single ports, ranges, and lists.
|
||||
// Examples: [80], [443], [80, 443], ["1024-65535"], [53, "80-90"]
|
||||
type PortSpec []string
|
||||
|
||||
func (ps *PortSpec) UnmarshalYAML(unmarshal func(interface{}) error) error {
|
||||
var multi []interface{}
|
||||
if err := unmarshal(&multi); err == nil {
|
||||
for _, v := range multi {
|
||||
switch val := v.(type) {
|
||||
case int:
|
||||
*ps = append(*ps, fmt.Sprintf("%d", val))
|
||||
case float64:
|
||||
*ps = append(*ps, fmt.Sprintf("%d", int(val)))
|
||||
case string:
|
||||
*ps = append(*ps, val)
|
||||
default:
|
||||
return fmt.Errorf("unsupported port value type %T", v)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var single string
|
||||
if err := unmarshal(&single); err == nil {
|
||||
*ps = PortSpec{single}
|
||||
return nil
|
||||
}
|
||||
|
||||
var num int
|
||||
if err := unmarshal(&num); err == nil {
|
||||
*ps = PortSpec{fmt.Sprintf("%d", num)}
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("invalid port spec")
|
||||
}
|
||||
|
||||
func (c *Config) validateRules() error {
|
||||
for i, r := range c.Rules {
|
||||
switch r.Action {
|
||||
case RuleAccept, RuleDrop, RuleReject, RuleDNAT, RuleRedirect, RuleLog:
|
||||
default:
|
||||
return fmt.Errorf("rule[%d]: unknown action %q", i, r.Action)
|
||||
}
|
||||
|
||||
if r.Source == "" {
|
||||
return fmt.Errorf("rule[%d]: source required", i)
|
||||
}
|
||||
if r.Dest == "" {
|
||||
return fmt.Errorf("rule[%d]: dest required", i)
|
||||
}
|
||||
|
||||
srcZone := zoneFromSpec(r.Source)
|
||||
if srcZone != "all" {
|
||||
if _, ok := c.Zones[srcZone]; !ok {
|
||||
return fmt.Errorf("rule[%d]: source zone %q not defined", i, srcZone)
|
||||
}
|
||||
}
|
||||
|
||||
dstZone := zoneFromSpec(r.Dest)
|
||||
if dstZone != "all" {
|
||||
if _, ok := c.Zones[dstZone]; !ok {
|
||||
return fmt.Errorf("rule[%d]: dest zone %q not defined", i, dstZone)
|
||||
}
|
||||
}
|
||||
|
||||
if r.PortGroup != "" {
|
||||
if _, ok := c.PortGroups[r.PortGroup]; !ok {
|
||||
return fmt.Errorf("rule[%d]: portgroup %q not defined", i, r.PortGroup)
|
||||
}
|
||||
if r.Proto != "" || len(r.DPort) > 0 {
|
||||
return fmt.Errorf("rule[%d]: portgroup is mutually exclusive with proto/dport", i)
|
||||
}
|
||||
}
|
||||
|
||||
if r.Action == RuleDNAT && r.DNATDest == "" {
|
||||
return fmt.Errorf("rule[%d]: dnat_dest required for DNAT action", i)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24".
|
||||
func zoneFromSpec(spec string) string {
|
||||
for i, c := range spec {
|
||||
if c == ':' {
|
||||
return spec[:i]
|
||||
}
|
||||
}
|
||||
return spec
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type SNATAction string
|
||||
|
||||
const (
|
||||
SNATMasquerade SNATAction = "masquerade"
|
||||
SNATAddress SNATAction = "snat"
|
||||
)
|
||||
|
||||
type SNATRule struct {
|
||||
Action SNATAction `yaml:"action"`
|
||||
Address string `yaml:"address,omitempty"`
|
||||
Source string `yaml:"source,omitempty"`
|
||||
DestInterface string `yaml:"dest_interface"`
|
||||
Proto string `yaml:"proto,omitempty"`
|
||||
DPort PortSpec `yaml:"dport,omitempty"`
|
||||
Comment string `yaml:"comment,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Config) validateSNAT() error {
|
||||
for i, s := range c.SNAT {
|
||||
switch s.Action {
|
||||
case SNATMasquerade, SNATAddress:
|
||||
default:
|
||||
return fmt.Errorf("snat[%d]: unknown action %q", i, s.Action)
|
||||
}
|
||||
|
||||
if s.Action == SNATAddress && s.Address == "" {
|
||||
return fmt.Errorf("snat[%d]: address required for snat action", i)
|
||||
}
|
||||
|
||||
if s.DestInterface == "" {
|
||||
return fmt.Errorf("snat[%d]: dest_interface required", i)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package config
|
||||
|
||||
import "fmt"
|
||||
|
||||
type ZoneType string
|
||||
|
||||
const (
|
||||
ZoneIP ZoneType = "ip"
|
||||
ZoneIPSec ZoneType = "ipsec"
|
||||
ZoneFirewall ZoneType = "firewall"
|
||||
ZoneLoopback ZoneType = "loopback"
|
||||
)
|
||||
|
||||
type Zone struct {
|
||||
Type ZoneType `yaml:"type"`
|
||||
Parent string `yaml:"parent,omitempty"`
|
||||
Options []string `yaml:"options,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Config) validateZones() error {
|
||||
if len(c.Zones) == 0 {
|
||||
return fmt.Errorf("no zones defined")
|
||||
}
|
||||
|
||||
firewallCount := 0
|
||||
for name, z := range c.Zones {
|
||||
switch z.Type {
|
||||
case ZoneIP, ZoneIPSec, ZoneFirewall, ZoneLoopback:
|
||||
default:
|
||||
return fmt.Errorf("zone %q: unknown type %q", name, z.Type)
|
||||
}
|
||||
if z.Type == ZoneFirewall {
|
||||
firewallCount++
|
||||
}
|
||||
if z.Parent != "" {
|
||||
if _, ok := c.Zones[z.Parent]; !ok {
|
||||
return fmt.Errorf("zone %q: parent zone %q not defined", name, z.Parent)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if firewallCount != 1 {
|
||||
return fmt.Errorf("exactly one firewall zone required, found %d", firewallCount)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/google/nftables"
|
||||
)
|
||||
|
||||
type ForeignRule struct {
|
||||
Table string
|
||||
Chain string
|
||||
Handle uint64
|
||||
Family nftables.TableFamily
|
||||
|
||||
table *nftables.Table
|
||||
chain *nftables.Chain
|
||||
}
|
||||
|
||||
func (f ForeignRule) String() string {
|
||||
return fmt.Sprintf("table=%s chain=%s handle=%d", f.Table, f.Chain, f.Handle)
|
||||
}
|
||||
|
||||
func (e *Engine) FindForeignRules() ([]ForeignRule, error) {
|
||||
tables, err := e.conn.ListTables()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listing tables: %w", err)
|
||||
}
|
||||
|
||||
var ourTable *nftables.Table
|
||||
for _, t := range tables {
|
||||
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet {
|
||||
ourTable = t
|
||||
break
|
||||
}
|
||||
}
|
||||
if ourTable == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
compiler := NewCompiler(e.cfg)
|
||||
desired, err := compiler.Compile()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("compiling config: %w", err)
|
||||
}
|
||||
|
||||
desiredTags := make(map[string]bool)
|
||||
for _, rules := range desired.Rules {
|
||||
for _, r := range rules {
|
||||
desiredTags[r.Tag] = true
|
||||
}
|
||||
}
|
||||
|
||||
var foreign []ForeignRule
|
||||
|
||||
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listing chains: %w", err)
|
||||
}
|
||||
|
||||
for _, chain := range chains {
|
||||
if chain.Table.Name != e.cfg.Settings.TableName {
|
||||
continue
|
||||
}
|
||||
rules, err := e.conn.GetRules(ourTable, chain)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, rule := range rules {
|
||||
tag := string(rule.UserData)
|
||||
if tag == "" || !desiredTags[tag] {
|
||||
foreign = append(foreign, ForeignRule{
|
||||
Table: ourTable.Name,
|
||||
Chain: chain.Name,
|
||||
Handle: rule.Handle,
|
||||
Family: ourTable.Family,
|
||||
table: ourTable,
|
||||
chain: chain,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return foreign, nil
|
||||
}
|
||||
|
||||
func (e *Engine) PurgeForeignRules(foreign []ForeignRule) error {
|
||||
for _, f := range foreign {
|
||||
e.conn.DelRule(&nftables.Rule{
|
||||
Table: f.table,
|
||||
Chain: f.chain,
|
||||
Handle: f.Handle,
|
||||
})
|
||||
}
|
||||
return e.conn.Flush()
|
||||
}
|
||||
@@ -0,0 +1,409 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/config"
|
||||
)
|
||||
|
||||
type Compiler struct {
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
func NewCompiler(cfg *config.Config) *Compiler {
|
||||
return &Compiler{cfg: cfg}
|
||||
}
|
||||
|
||||
func (c *Compiler) Compile() (*FirewallState, error) {
|
||||
state := &FirewallState{
|
||||
Rules: make(map[string][]ManagedRule),
|
||||
}
|
||||
|
||||
if err := c.compileRules(state); err != nil {
|
||||
return nil, fmt.Errorf("rules: %w", err)
|
||||
}
|
||||
if err := c.compilePolicies(state); err != nil {
|
||||
return nil, fmt.Errorf("policies: %w", err)
|
||||
}
|
||||
if err := c.compileSNAT(state); err != nil {
|
||||
return nil, fmt.Errorf("snat: %w", err)
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (c *Compiler) compileRules(state *FirewallState) error {
|
||||
fwZone := c.cfg.FirewallZone()
|
||||
|
||||
for i, rule := range c.cfg.Rules {
|
||||
tag := fmt.Sprintf("rule:%d", i)
|
||||
|
||||
proto := rule.Proto
|
||||
var ports config.PortSpec
|
||||
if rule.PortGroup != "" {
|
||||
pg, _ := c.cfg.ResolvePortGroup(rule.PortGroup)
|
||||
proto = pg.Proto
|
||||
ports = pg.Ports
|
||||
} else {
|
||||
ports = rule.DPort
|
||||
}
|
||||
|
||||
srcZone, _ := splitZoneSpec(rule.Source)
|
||||
dstZone, _ := splitZoneSpec(rule.Dest)
|
||||
|
||||
srcIfaces := c.resolveZoneInterfaces(srcZone)
|
||||
dstIfaces := c.resolveZoneInterfaces(dstZone)
|
||||
|
||||
chain := c.selectChain(srcZone, dstZone, fwZone)
|
||||
|
||||
for _, srcIface := range srcIfaces {
|
||||
for _, dstIface := range dstIfaces {
|
||||
exprs, err := c.buildRuleExprs(srcIface, dstIface, chain, proto, ports, rule.Action, rule.Source, rule.Dest)
|
||||
if err != nil {
|
||||
return fmt.Errorf("rule[%d]: %w", i, err)
|
||||
}
|
||||
state.Rules[chain] = append(state.Rules[chain], ManagedRule{
|
||||
Chain: chain,
|
||||
Exprs: exprs,
|
||||
Tag: tag,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Compiler) compilePolicies(state *FirewallState) error {
|
||||
fwZone := c.cfg.FirewallZone()
|
||||
|
||||
for i, pol := range c.cfg.Policy {
|
||||
tag := fmt.Sprintf("policy:%d", i)
|
||||
|
||||
srcZones := c.expandZoneRef(pol.Source)
|
||||
dstZones := c.expandZoneRef(pol.Dest)
|
||||
|
||||
for _, sz := range srcZones {
|
||||
for _, dz := range dstZones {
|
||||
if sz == dz {
|
||||
continue
|
||||
}
|
||||
|
||||
chain := c.selectChain(sz, dz, fwZone)
|
||||
srcIfaces := c.resolveZoneInterfaces(sz)
|
||||
dstIfaces := c.resolveZoneInterfaces(dz)
|
||||
|
||||
for _, si := range srcIfaces {
|
||||
for _, di := range dstIfaces {
|
||||
var exprs []expr.Any
|
||||
|
||||
if si != "" {
|
||||
exprs = append(exprs, matchIface(true, si)...)
|
||||
}
|
||||
if di != "" && chain == "forward" {
|
||||
exprs = append(exprs, matchIface(false, di)...)
|
||||
}
|
||||
|
||||
exprs = append(exprs, policyVerdict(pol.Action)...)
|
||||
|
||||
state.Rules[chain] = append(state.Rules[chain], ManagedRule{
|
||||
Chain: chain,
|
||||
Exprs: exprs,
|
||||
Tag: tag,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Compiler) compileSNAT(state *FirewallState) error {
|
||||
for i, snat := range c.cfg.SNAT {
|
||||
tag := fmt.Sprintf("snat:%d", i)
|
||||
|
||||
var exprs []expr.Any
|
||||
|
||||
exprs = append(exprs, matchIface(false, snat.DestInterface)...)
|
||||
|
||||
if snat.Source != "" {
|
||||
srcExprs, err := matchSourceCIDR(snat.Source)
|
||||
if err != nil {
|
||||
return fmt.Errorf("snat[%d]: %w", i, err)
|
||||
}
|
||||
exprs = append(exprs, srcExprs...)
|
||||
}
|
||||
|
||||
if snat.Proto != "" {
|
||||
exprs = append(exprs, matchProto(snat.Proto)...)
|
||||
}
|
||||
|
||||
switch snat.Action {
|
||||
case config.SNATMasquerade:
|
||||
exprs = append(exprs, &expr.Masq{})
|
||||
case config.SNATAddress:
|
||||
ip := net.ParseIP(snat.Address)
|
||||
if ip == nil {
|
||||
return fmt.Errorf("snat[%d]: invalid address %q", i, snat.Address)
|
||||
}
|
||||
ip4 := ip.To4()
|
||||
if ip4 != nil {
|
||||
exprs = append(exprs,
|
||||
&expr.Immediate{Register: 1, Data: ip4},
|
||||
&expr.NAT{
|
||||
Type: expr.NATTypeSourceNAT,
|
||||
Family: unix.NFPROTO_IPV4,
|
||||
RegAddrMin: 1,
|
||||
RegAddrMax: 1,
|
||||
},
|
||||
)
|
||||
} else {
|
||||
exprs = append(exprs,
|
||||
&expr.Immediate{Register: 1, Data: ip.To16()},
|
||||
&expr.NAT{
|
||||
Type: expr.NATTypeSourceNAT,
|
||||
Family: unix.NFPROTO_IPV6,
|
||||
RegAddrMin: 1,
|
||||
RegAddrMax: 1,
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{
|
||||
Chain: "postrouting",
|
||||
Exprs: exprs,
|
||||
Tag: tag,
|
||||
})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string {
|
||||
if dstZone == fwZone {
|
||||
return "input"
|
||||
}
|
||||
if srcZone == fwZone {
|
||||
return "output"
|
||||
}
|
||||
return "forward"
|
||||
}
|
||||
|
||||
func (c *Compiler) resolveZoneInterfaces(zone string) []string {
|
||||
if zone == "all" || zone == "" {
|
||||
return []string{""}
|
||||
}
|
||||
ifaces := c.cfg.ZoneInterfaces(zone)
|
||||
if len(ifaces) == 0 {
|
||||
return []string{""}
|
||||
}
|
||||
return ifaces
|
||||
}
|
||||
|
||||
func (c *Compiler) expandZoneRef(ref string) []string {
|
||||
if ref == "all" {
|
||||
var zones []string
|
||||
for name := range c.cfg.Zones {
|
||||
zones = append(zones, name)
|
||||
}
|
||||
return zones
|
||||
}
|
||||
return []string{ref}
|
||||
}
|
||||
|
||||
func (c *Compiler) buildRuleExprs(srcIface, dstIface, chain, proto string, ports config.PortSpec, action config.RuleAction, srcSpec, dstSpec string) ([]expr.Any, error) {
|
||||
var exprs []expr.Any
|
||||
|
||||
if srcIface != "" {
|
||||
exprs = append(exprs, matchIface(true, srcIface)...)
|
||||
}
|
||||
if dstIface != "" && chain == "forward" {
|
||||
exprs = append(exprs, matchIface(false, dstIface)...)
|
||||
}
|
||||
|
||||
_, srcAddr := splitZoneSpec(srcSpec)
|
||||
_, dstAddr := splitZoneSpec(dstSpec)
|
||||
|
||||
if srcAddr != "" {
|
||||
src, err := matchSourceCIDR(srcAddr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
exprs = append(exprs, src...)
|
||||
}
|
||||
|
||||
if dstAddr != "" {
|
||||
dst, err := matchDestCIDR(dstAddr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
exprs = append(exprs, dst...)
|
||||
}
|
||||
|
||||
if proto != "" {
|
||||
exprs = append(exprs, matchProto(proto)...)
|
||||
}
|
||||
|
||||
for _, portStr := range ports {
|
||||
p, err := parsePort(portStr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
exprs = append(exprs, matchDPort(p)...)
|
||||
}
|
||||
|
||||
switch action {
|
||||
case config.RuleAccept:
|
||||
exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictAccept})
|
||||
case config.RuleDrop:
|
||||
exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop})
|
||||
case config.RuleReject:
|
||||
exprs = append(exprs, &expr.Reject{})
|
||||
}
|
||||
|
||||
return exprs, nil
|
||||
}
|
||||
|
||||
func matchIface(input bool, name string) []expr.Any {
|
||||
key := expr.MetaKeyOIFNAME
|
||||
if input {
|
||||
key = expr.MetaKeyIIFNAME
|
||||
}
|
||||
padded := make([]byte, 16)
|
||||
copy(padded, name+"\x00")
|
||||
return []expr.Any{
|
||||
&expr.Meta{Key: key, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: padded[:len(name)+1]},
|
||||
}
|
||||
}
|
||||
|
||||
func matchProto(proto string) []expr.Any {
|
||||
var protoNum byte
|
||||
switch strings.ToLower(proto) {
|
||||
case "tcp":
|
||||
protoNum = unix.IPPROTO_TCP
|
||||
case "udp":
|
||||
protoNum = unix.IPPROTO_UDP
|
||||
case "icmp":
|
||||
protoNum = unix.IPPROTO_ICMP
|
||||
default:
|
||||
n, _ := strconv.Atoi(proto)
|
||||
protoNum = byte(n)
|
||||
}
|
||||
return []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}},
|
||||
}
|
||||
}
|
||||
|
||||
func matchDPort(port uint16) []expr.Any {
|
||||
portBytes := make([]byte, 2)
|
||||
binary.BigEndian.PutUint16(portBytes, port)
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: portBytes},
|
||||
}
|
||||
}
|
||||
|
||||
func matchSourceCIDR(cidr string) ([]expr.Any, error) {
|
||||
ip, ipNet, err := net.ParseCIDR(cidr)
|
||||
if err != nil {
|
||||
singleIP := net.ParseIP(cidr)
|
||||
if singleIP == nil {
|
||||
return nil, fmt.Errorf("invalid source address %q", cidr)
|
||||
}
|
||||
ip4 := singleIP.To4()
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ip4},
|
||||
}, nil
|
||||
}
|
||||
|
||||
ip4 := ip.To4()
|
||||
if ip4 == nil {
|
||||
return nil, fmt.Errorf("IPv6 source addresses not yet supported: %s", cidr)
|
||||
}
|
||||
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
|
||||
&expr.Bitwise{
|
||||
SourceRegister: 1,
|
||||
DestRegister: 1,
|
||||
Len: 4,
|
||||
Mask: ipNet.Mask,
|
||||
Xor: []byte{0, 0, 0, 0},
|
||||
},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ipNet.IP.To4()},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func matchDestCIDR(cidr string) ([]expr.Any, error) {
|
||||
ip, ipNet, err := net.ParseCIDR(cidr)
|
||||
if err != nil {
|
||||
singleIP := net.ParseIP(cidr)
|
||||
if singleIP == nil {
|
||||
return nil, fmt.Errorf("invalid dest address %q", cidr)
|
||||
}
|
||||
ip4 := singleIP.To4()
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ip4},
|
||||
}, nil
|
||||
}
|
||||
|
||||
ip4 := ip.To4()
|
||||
if ip4 == nil {
|
||||
return nil, fmt.Errorf("IPv6 dest addresses not yet supported: %s", cidr)
|
||||
}
|
||||
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4},
|
||||
&expr.Bitwise{
|
||||
SourceRegister: 1,
|
||||
DestRegister: 1,
|
||||
Len: 4,
|
||||
Mask: ipNet.Mask,
|
||||
Xor: []byte{0, 0, 0, 0},
|
||||
},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ipNet.IP.To4()},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func policyVerdict(action config.PolicyAction) []expr.Any {
|
||||
switch action {
|
||||
case config.PolicyAccept:
|
||||
return []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}}
|
||||
case config.PolicyDrop:
|
||||
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
|
||||
case config.PolicyReject:
|
||||
return []expr.Any{&expr.Reject{}}
|
||||
default:
|
||||
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
|
||||
}
|
||||
}
|
||||
|
||||
func splitZoneSpec(spec string) (zone, addr string) {
|
||||
idx := strings.IndexByte(spec, ':')
|
||||
if idx < 0 {
|
||||
return spec, ""
|
||||
}
|
||||
return spec[:idx], spec[idx+1:]
|
||||
}
|
||||
|
||||
func parsePort(s string) (uint16, error) {
|
||||
n, err := strconv.ParseUint(s, 10, 16)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid port %q: %w", s, err)
|
||||
}
|
||||
return uint16(n), nil
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
)
|
||||
|
||||
type ManagedRule struct {
|
||||
Chain string
|
||||
Handle uint64
|
||||
Exprs []expr.Any
|
||||
Tag string
|
||||
}
|
||||
|
||||
type FirewallState struct {
|
||||
Rules map[string][]ManagedRule
|
||||
}
|
||||
|
||||
type ChangeSet struct {
|
||||
Add []ManagedRule
|
||||
Remove []ManagedRule
|
||||
}
|
||||
|
||||
func (cs *ChangeSet) Empty() bool {
|
||||
return len(cs.Add) == 0 && len(cs.Remove) == 0
|
||||
}
|
||||
|
||||
func (cs *ChangeSet) Summary() string {
|
||||
var b strings.Builder
|
||||
if len(cs.Add) > 0 {
|
||||
fmt.Fprintf(&b, " + %d rule(s) to add\n", len(cs.Add))
|
||||
for _, r := range cs.Add {
|
||||
fmt.Fprintf(&b, " + [%s] %s\n", r.Chain, r.Tag)
|
||||
}
|
||||
}
|
||||
if len(cs.Remove) > 0 {
|
||||
fmt.Fprintf(&b, " - %d rule(s) to remove\n", len(cs.Remove))
|
||||
for _, r := range cs.Remove {
|
||||
fmt.Fprintf(&b, " - [%s] %s (handle %d)\n", r.Chain, r.Tag, r.Handle)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func computeDiff(current, desired *FirewallState) *ChangeSet {
|
||||
cs := &ChangeSet{}
|
||||
|
||||
currentByTag := make(map[string][]ManagedRule)
|
||||
for _, rules := range current.Rules {
|
||||
for _, r := range rules {
|
||||
if r.Tag != "" {
|
||||
currentByTag[r.Tag] = append(currentByTag[r.Tag], r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
desiredByTag := make(map[string][]ManagedRule)
|
||||
for _, rules := range desired.Rules {
|
||||
for _, r := range rules {
|
||||
desiredByTag[r.Tag] = append(desiredByTag[r.Tag], r)
|
||||
}
|
||||
}
|
||||
|
||||
for tag, desiredRules := range desiredByTag {
|
||||
if _, exists := currentByTag[tag]; !exists {
|
||||
cs.Add = append(cs.Add, desiredRules...)
|
||||
}
|
||||
}
|
||||
|
||||
for tag, currentRules := range currentByTag {
|
||||
if _, exists := desiredByTag[tag]; !exists {
|
||||
cs.Remove = append(cs.Remove, currentRules...)
|
||||
}
|
||||
}
|
||||
|
||||
return cs
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/google/nftables"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/config"
|
||||
)
|
||||
|
||||
type Engine struct {
|
||||
cfg *config.Config
|
||||
conn *nftables.Conn
|
||||
}
|
||||
|
||||
func NewEngine(cfg *config.Config) (*Engine, error) {
|
||||
conn, err := nftables.New()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connecting to nftables: %w", err)
|
||||
}
|
||||
return &Engine{cfg: cfg, conn: conn}, nil
|
||||
}
|
||||
|
||||
func (e *Engine) ensureTable() *nftables.Table {
|
||||
return e.conn.AddTable(&nftables.Table{
|
||||
Family: nftables.TableFamilyINet,
|
||||
Name: e.cfg.Settings.TableName,
|
||||
})
|
||||
}
|
||||
|
||||
func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain {
|
||||
chains := map[string]*nftables.Chain{
|
||||
"input": {
|
||||
Name: "input",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookInput,
|
||||
Priority: nftables.ChainPriorityFilter,
|
||||
Policy: policyPtr(nftables.ChainPolicyDrop),
|
||||
},
|
||||
"forward": {
|
||||
Name: "forward",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookForward,
|
||||
Priority: nftables.ChainPriorityFilter,
|
||||
Policy: policyPtr(nftables.ChainPolicyDrop),
|
||||
},
|
||||
"output": {
|
||||
Name: "output",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookOutput,
|
||||
Priority: nftables.ChainPriorityFilter,
|
||||
Policy: policyPtr(nftables.ChainPolicyAccept),
|
||||
},
|
||||
"postrouting": {
|
||||
Name: "postrouting",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeNAT,
|
||||
Hooknum: nftables.ChainHookPostrouting,
|
||||
Priority: nftables.ChainPriorityNATSource,
|
||||
},
|
||||
"prerouting": {
|
||||
Name: "prerouting",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeNAT,
|
||||
Hooknum: nftables.ChainHookPrerouting,
|
||||
Priority: nftables.ChainPriorityNATDest,
|
||||
},
|
||||
}
|
||||
|
||||
for name, chain := range chains {
|
||||
chains[name] = e.conn.AddChain(chain)
|
||||
}
|
||||
return chains
|
||||
}
|
||||
|
||||
func (e *Engine) Plan() (*ChangeSet, error) {
|
||||
compiler := NewCompiler(e.cfg)
|
||||
|
||||
desired, err := compiler.Compile()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("compiling config: %w", err)
|
||||
}
|
||||
|
||||
current, err := e.readCurrentState()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading current state: %w", err)
|
||||
}
|
||||
|
||||
return computeDiff(current, desired), nil
|
||||
}
|
||||
|
||||
func (e *Engine) Apply(changes *ChangeSet) error {
|
||||
table := e.ensureTable()
|
||||
chains := e.ensureChains(table)
|
||||
|
||||
for _, r := range changes.Remove {
|
||||
e.conn.DelRule(&nftables.Rule{
|
||||
Table: table,
|
||||
Chain: chains[r.Chain],
|
||||
Handle: r.Handle,
|
||||
})
|
||||
}
|
||||
|
||||
for _, r := range changes.Add {
|
||||
chain, ok := chains[r.Chain]
|
||||
if !ok {
|
||||
return fmt.Errorf("unknown chain %q", r.Chain)
|
||||
}
|
||||
e.conn.AddRule(&nftables.Rule{
|
||||
Table: table,
|
||||
Chain: chain,
|
||||
Exprs: r.Exprs,
|
||||
UserData: []byte(r.Tag),
|
||||
})
|
||||
}
|
||||
|
||||
return e.conn.Flush()
|
||||
}
|
||||
|
||||
func (e *Engine) Flush() error {
|
||||
tables, err := e.conn.ListTables()
|
||||
if err != nil {
|
||||
return fmt.Errorf("listing tables: %w", err)
|
||||
}
|
||||
|
||||
for _, t := range tables {
|
||||
if t.Name == e.cfg.Settings.TableName {
|
||||
e.conn.DelTable(t)
|
||||
return e.conn.Flush()
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Engine) readCurrentState() (*FirewallState, error) {
|
||||
state := &FirewallState{
|
||||
Rules: make(map[string][]ManagedRule),
|
||||
}
|
||||
|
||||
tables, err := e.conn.ListTables()
|
||||
if err != nil {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
var ourTable *nftables.Table
|
||||
for _, t := range tables {
|
||||
if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet {
|
||||
ourTable = t
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if ourTable == nil {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
|
||||
if err != nil {
|
||||
return state, nil
|
||||
}
|
||||
|
||||
for _, chain := range chains {
|
||||
if chain.Table.Name != e.cfg.Settings.TableName {
|
||||
continue
|
||||
}
|
||||
rules, err := e.conn.GetRules(ourTable, chain)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, rule := range rules {
|
||||
state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{
|
||||
Chain: chain.Name,
|
||||
Handle: rule.Handle,
|
||||
Exprs: rule.Exprs,
|
||||
Tag: string(rule.UserData),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy {
|
||||
return &p
|
||||
}
|
||||
Reference in New Issue
Block a user