195 lines
4.1 KiB
Go
195 lines
4.1 KiB
Go
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)
|
|
}
|
|
rule := &nftables.Rule{
|
|
Table: table,
|
|
Chain: chain,
|
|
Exprs: r.Exprs,
|
|
UserData: []byte(r.Tag),
|
|
}
|
|
if r.Before != 0 {
|
|
rule.Position = r.Before
|
|
e.conn.InsertRule(rule)
|
|
} else {
|
|
e.conn.AddRule(rule)
|
|
}
|
|
}
|
|
|
|
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
|
|
}
|