mirror of
https://github.com/qdm12/gluetun.git
synced 2026-08-11 14:52:56 +02:00
wip ai
This commit is contained in:
@@ -48,3 +48,15 @@ func (f *Firewall) SetBaseChainsPolicy(_ context.Context, policy string) error {
|
|||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetIPv4AllPolicies sets the policy of all base chains to ACCEPT or DROP.
|
||||||
|
// In nftables with inet family, this also affects IPv6 rules.
|
||||||
|
func (f *Firewall) SetIPv4AllPolicies(ctx context.Context, policy string) error {
|
||||||
|
return f.SetBaseChainsPolicy(ctx, policy)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetIPv6AllPolicies sets the policy of all base chains to ACCEPT or DROP.
|
||||||
|
// In nftables with inet family, this also affects IPv4 rules.
|
||||||
|
func (f *Firewall) SetIPv6AllPolicies(ctx context.Context, policy string) error {
|
||||||
|
return f.SetBaseChainsPolicy(ctx, policy)
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,9 +3,11 @@ package nftables
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
"github.com/google/nftables"
|
"github.com/google/nftables"
|
||||||
"github.com/google/nftables/expr"
|
"github.com/google/nftables/expr"
|
||||||
|
"github.com/qdm12/gluetun/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Firewall) AcceptIpv6MulticastOutput(_ context.Context, intf string) error {
|
func (f *Firewall) AcceptIpv6MulticastOutput(_ context.Context, intf string) error {
|
||||||
@@ -76,3 +78,576 @@ func (f *Firewall) AcceptIpv6MulticastOutput(_ context.Context, intf string) err
|
|||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) AcceptOutputTrafficToVPN(_ context.Context, defaultInterface string,
|
||||||
|
connection models.Connection, remove bool) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
// Prepare the destination IP and port
|
||||||
|
const maxExprsLen = 7
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
// Interface filter
|
||||||
|
if defaultInterface != "" && defaultInterface != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(defaultInterface + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Destination IP address
|
||||||
|
if connection.IP.Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 16, // IPv4 destination address offset
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: connection.IP.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else { // IPv6
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 24, // IPv6 destination address offset
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: connection.IP.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Protocol (tcp or udp)
|
||||||
|
var protocolByte uint8
|
||||||
|
if connection.Protocol == "tcp" || connection.Protocol == "tcp-client" {
|
||||||
|
protocolByte = 6 // TCP
|
||||||
|
} else if connection.Protocol == "udp" {
|
||||||
|
protocolByte = 17 // UDP
|
||||||
|
} else {
|
||||||
|
return fmt.Errorf("unsupported protocol: %s", connection.Protocol)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 3, // Protocol byte offset in IP header
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protocolByte},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Destination port
|
||||||
|
portBytes := []byte{byte(connection.Port >> 8), byte(connection.Port)} //nolint:mnd
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2, // destination port offset
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: portBytes,
|
||||||
|
},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, rule)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
if !remove {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
}
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) AcceptOutput(_ context.Context, protocol, intf string, ip netip.Addr, port uint16, remove bool) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
const maxExprsLen = 7
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
if intf != "" && intf != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(intf + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip.Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 16, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ip.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 24, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: ip.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
var protocolByte uint8
|
||||||
|
switch protocol {
|
||||||
|
case "tcp":
|
||||||
|
protocolByte = 6 //nolint:mnd
|
||||||
|
case "udp":
|
||||||
|
protocolByte = 17 //nolint:mnd
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported protocol: %s", protocol)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 3, //nolint:mnd
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protocolByte},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
portBytes := []byte{byte(port >> 8), byte(port)} //nolint:mnd
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2, //nolint:mnd
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: portBytes,
|
||||||
|
},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, rule)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
if !remove {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
}
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) AcceptOutputFromIPPortToIPPort(_ context.Context, protocol, intf string,
|
||||||
|
source, destination netip.AddrPort, remove bool,
|
||||||
|
) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
const maxExprsLen = 10 //nolint:mnd
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
if intf != "" && intf != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(intf + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if source.Addr().Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 12, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: source.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 8, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: source.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if destination.Addr().Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 16, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: destination.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 24, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: destination.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
var protocolByte uint8
|
||||||
|
switch protocol {
|
||||||
|
case "tcp":
|
||||||
|
protocolByte = 6 //nolint:mnd
|
||||||
|
case "udp":
|
||||||
|
protocolByte = 17 //nolint:mnd
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported protocol: %s", protocol)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 3, //nolint:mnd
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{protocolByte},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
sourcePortBytes := []byte{byte(source.Port() >> 8), byte(source.Port())} //nolint:mnd
|
||||||
|
destinationPortBytes := []byte{byte(destination.Port() >> 8), byte(destination.Port())} //nolint:mnd
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 0, //nolint:mnd
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: sourcePortBytes,
|
||||||
|
},
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2, //nolint:mnd
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: destinationPortBytes,
|
||||||
|
},
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, rule)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
if !remove {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
}
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) AcceptOutputFromIPToSubnet(_ context.Context, intf string, assignedIP netip.Addr,
|
||||||
|
subnet netip.Prefix, remove bool,
|
||||||
|
) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
const maxExprsLen = 8
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
if intf != "" && intf != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(intf + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if assignedIP.Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 12, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: assignedIP.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 8, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: assignedIP.AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if subnet.Addr().Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 16, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Bitwise{
|
||||||
|
SourceRegister: 1,
|
||||||
|
DestRegister: 1,
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
Mask: subnet.Masked().Addr().AsSlice(),
|
||||||
|
Xor: []byte{0, 0, 0, 0}, //nolint:mnd
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 24, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Bitwise{
|
||||||
|
SourceRegister: 1,
|
||||||
|
DestRegister: 1,
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
Mask: subnet.Masked().Addr().AsSlice(),
|
||||||
|
Xor: []byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, //nolint:mnd
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, rule)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
if !remove {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
}
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Firewall) AcceptOutputThroughInterface(_ context.Context, intf string, remove bool) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
const maxExprsLen = 3
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
if intf != "" && intf != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(intf + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||||
|
)
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, rule)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
if !remove {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
}
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,166 @@
|
|||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RedirectPort redirects incoming traffic on the specified source port to the
|
||||||
|
// specified destination port, for both TCP and UDP protocols, on the interface intf.
|
||||||
|
// If intf is empty or "*", the interface is not used as a filter. If remove is true,
|
||||||
|
// the redirection is removed instead of added. This is used for VPN server side
|
||||||
|
// port forwarding, with intf set to the VPN tunnel interface.
|
||||||
|
func (f *Firewall) RedirectPort(_ context.Context, intf string,
|
||||||
|
sourcePort, destinationPort uint16, remove bool,
|
||||||
|
) (err error) {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, inputChain, _, _ := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
natTable := conn.AddTable(&nftables.Table{
|
||||||
|
Family: nftables.TableFamilyINet,
|
||||||
|
Name: "nat",
|
||||||
|
})
|
||||||
|
|
||||||
|
preroutingChain := conn.AddChain(&nftables.Chain{
|
||||||
|
Name: "prerouting",
|
||||||
|
Table: natTable,
|
||||||
|
Type: nftables.ChainTypeNAT,
|
||||||
|
Hooknum: nftables.ChainHookPrerouting,
|
||||||
|
Priority: nftables.ChainPriorityNATDest,
|
||||||
|
})
|
||||||
|
|
||||||
|
sourcePortBytes := []byte{byte(sourcePort >> 8), byte(sourcePort)} //nolint:mnd
|
||||||
|
destinationPortBytes := []byte{byte(destinationPort >> 8), byte(destinationPort)} //nolint:mnd
|
||||||
|
const tcp, udp uint8 = 6, 17 //nolint:mnd
|
||||||
|
protocols := []uint8{tcp, udp}
|
||||||
|
|
||||||
|
var rulesToDelete []*nftables.Rule
|
||||||
|
|
||||||
|
for _, protocol := range protocols {
|
||||||
|
prerouteRule := buildRedirectRule(conn, natTable, preroutingChain,
|
||||||
|
intf, protocol, sourcePortBytes, destinationPort)
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(prerouteRule)
|
||||||
|
f.rules = append(f.rules, prerouteRule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, prerouteRule)
|
||||||
|
if err != nil {
|
||||||
|
rulesToDelete = append(rulesToDelete, prerouteRule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
inputRule := buildRedirectInputRule(table, inputChain,
|
||||||
|
intf, protocol, destinationPortBytes)
|
||||||
|
|
||||||
|
if !remove {
|
||||||
|
conn.AddRule(inputRule)
|
||||||
|
f.rules = append(f.rules, inputRule)
|
||||||
|
} else {
|
||||||
|
err = f.deleteRule(conn, inputRule)
|
||||||
|
if err != nil {
|
||||||
|
rulesToDelete = append(rulesToDelete, inputRule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil && !isTableDoesNotExist(err) {
|
||||||
|
if !remove {
|
||||||
|
removeFailedRules(f.rules, rulesToDelete)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("redirecting source port %d to destination port %d on interface %s: %w",
|
||||||
|
sourcePort, destinationPort, intf, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if isTableDoesNotExist(err) && !remove {
|
||||||
|
f.logger.Warnf("IPv6 port redirection disabled because your kernel does not support IPv6 NAT: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildRedirectRule(_ *nftables.Conn, natTable *nftables.Table,
|
||||||
|
preroutingChain *nftables.Chain, intf string, protocol uint8,
|
||||||
|
sourcePortBytes []byte, destinationPort uint16,
|
||||||
|
) *nftables.Rule {
|
||||||
|
const regProto uint32 = 2
|
||||||
|
portReg := make([]byte, regProto)
|
||||||
|
binary.BigEndian.PutUint16(portReg, destinationPort)
|
||||||
|
|
||||||
|
exprs := buildRedirectMatchExprs(intf, protocol, sourcePortBytes)
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Immediate{Register: regProto, Data: portReg},
|
||||||
|
&expr.NAT{
|
||||||
|
Type: expr.NATTypeDestNAT,
|
||||||
|
Family: uint32(nftables.TableFamilyINet),
|
||||||
|
RegProtoMin: regProto,
|
||||||
|
RegProtoMax: regProto,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return &nftables.Rule{
|
||||||
|
Table: natTable,
|
||||||
|
Chain: preroutingChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildRedirectInputRule(table *nftables.Table, inputChain *nftables.Chain,
|
||||||
|
intf string, protocol uint8, destinationPortBytes []byte,
|
||||||
|
) *nftables.Rule {
|
||||||
|
exprs := buildRedirectMatchExprs(intf, protocol, destinationPortBytes)
|
||||||
|
exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictAccept})
|
||||||
|
|
||||||
|
return &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: inputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildRedirectMatchExprs(intf string, protocol uint8, portBytes []byte) []expr.Any {
|
||||||
|
const maxExprsLen = 6
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
if intf != "" && intf != "*" {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte(intf + "\x00")},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 9, Len: 1}, //nolint:mnd
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protocol}},
|
||||||
|
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, //nolint:mnd
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: portBytes},
|
||||||
|
)
|
||||||
|
|
||||||
|
return exprs
|
||||||
|
}
|
||||||
|
|
||||||
|
func isTableDoesNotExist(err error) bool {
|
||||||
|
return strings.Contains(err.Error(), "Table does not exist")
|
||||||
|
}
|
||||||
|
|
||||||
|
func removeFailedRules(rules []*nftables.Rule, failed []*nftables.Rule) {
|
||||||
|
for i := len(rules) - 1; i >= 0; i-- {
|
||||||
|
if slices.Contains(failed, rules[i]) {
|
||||||
|
rules = append(rules[:i], rules[i+1:]...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,23 @@
|
|||||||
package nftables
|
package nftables
|
||||||
|
|
||||||
import "github.com/google/nftables"
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
iptablesCommand = "iptables-nft"
|
||||||
|
iptablesFallbackCmd = "iptables"
|
||||||
|
ip6tablesCommand = "ip6tables-nft"
|
||||||
|
ip6tablesFallbackCmd = "ip6tables"
|
||||||
|
)
|
||||||
|
|
||||||
func IsSupported() bool {
|
func IsSupported() bool {
|
||||||
conn, err := nftables.New()
|
conn, err := nftables.New()
|
||||||
@@ -10,3 +27,122 @@ func IsSupported() bool {
|
|||||||
_, err = conn.ListTable("filter")
|
_, err = conn.ListTable("filter")
|
||||||
return err == nil
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Version obtains the version of the installed nftables.
|
||||||
|
func (f *Firewall) Version(ctx context.Context) (string, error) {
|
||||||
|
const emptyVersionError = "nft version string is empty"
|
||||||
|
cmd := exec.CommandContext(ctx, "nft", "-v")
|
||||||
|
output, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("running nft -v: %w", err)
|
||||||
|
}
|
||||||
|
outputStr := strings.TrimSpace(string(output))
|
||||||
|
words := strings.Fields(outputStr)
|
||||||
|
if len(words) == 0 {
|
||||||
|
return "", errors.New(emptyVersionError) //nolint:err113
|
||||||
|
}
|
||||||
|
return words[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// findIptablesCommand finds the available iptables-nft or iptables command.
|
||||||
|
func findIptablesCommand() (string, error) {
|
||||||
|
if path, err := exec.LookPath(iptablesCommand); err == nil {
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
if path, err := exec.LookPath(iptablesFallbackCmd); err == nil {
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("iptables command not found: %s or %s", iptablesCommand, iptablesFallbackCmd) //nolint:err113
|
||||||
|
}
|
||||||
|
|
||||||
|
// findIP6tablesCommand finds the available ip6tables-nft or ip6tables command.
|
||||||
|
func findIP6tablesCommand() (string, error) {
|
||||||
|
if path, err := exec.LookPath(ip6tablesCommand); err == nil {
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
if path, err := exec.LookPath(ip6tablesFallbackCmd); err == nil {
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("ip6tables command not found: %s or %s", ip6tablesCommand, ip6tablesFallbackCmd) //nolint:err113
|
||||||
|
}
|
||||||
|
|
||||||
|
// RunUserPostRules reads and executes custom iptables-style rules from a file.
|
||||||
|
// Since iptables-nft is nftables under the hood, we delegate to it for rule
|
||||||
|
// parsing compatibility with user-written iptables rules.
|
||||||
|
func (f *Firewall) RunUserPostRules(ctx context.Context, filepath string) error {
|
||||||
|
file, err := os.OpenFile(filepath, os.O_RDONLY, 0)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return nil
|
||||||
|
} else if err != nil {
|
||||||
|
return fmt.Errorf("opening user rules file: %w", err)
|
||||||
|
}
|
||||||
|
content, err := io.ReadAll(file)
|
||||||
|
if err != nil {
|
||||||
|
_ = file.Close()
|
||||||
|
return fmt.Errorf("reading user rules file: %w", err)
|
||||||
|
}
|
||||||
|
if err := file.Close(); err != nil {
|
||||||
|
return fmt.Errorf("closing user rules file: %w", err)
|
||||||
|
}
|
||||||
|
lines := strings.Split(string(content), "\n")
|
||||||
|
|
||||||
|
iptablesCmd, err := findIptablesCommand()
|
||||||
|
if err != nil {
|
||||||
|
f.logger.Warnf("iptables-nft not available, skipping user post-rules for IPv4")
|
||||||
|
}
|
||||||
|
ip6tablesCmd, err := findIP6tablesCommand()
|
||||||
|
if err != nil {
|
||||||
|
f.logger.Warnf("ip6tables-nft not available, IPv6 user post-rules will fail")
|
||||||
|
}
|
||||||
|
|
||||||
|
for lineNum, line := range lines {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line == "" || strings.HasPrefix(line, "#") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var cmdName string
|
||||||
|
var ruleArgs string
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(line, "iptables "):
|
||||||
|
cmdName = iptablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "iptables ")
|
||||||
|
case strings.HasPrefix(line, "iptables-nft "):
|
||||||
|
cmdName = iptablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "iptables-nft ")
|
||||||
|
case strings.HasPrefix(line, "iptables-legacy "):
|
||||||
|
cmdName = iptablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "iptables-legacy ")
|
||||||
|
case strings.HasPrefix(line, "ip6tables "):
|
||||||
|
cmdName = ip6tablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "ip6tables ")
|
||||||
|
case strings.HasPrefix(line, "ip6tables-nft "):
|
||||||
|
cmdName = ip6tablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "ip6tables-nft ")
|
||||||
|
case strings.HasPrefix(line, "ip6tables-legacy "):
|
||||||
|
cmdName = ip6tablesCmd
|
||||||
|
ruleArgs = strings.TrimPrefix(line, "ip6tables-legacy ")
|
||||||
|
default:
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if cmdName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
args := strings.Fields(ruleArgs)
|
||||||
|
if len(args) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd := exec.CommandContext(ctx, cmdName, args...)
|
||||||
|
output, err := cmd.CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
outputStr := strings.TrimSpace(string(output))
|
||||||
|
return fmt.Errorf("running user rule on line %d (%s %s): %w: %s",
|
||||||
|
lineNum+1, cmdName, ruleArgs, err, outputStr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,207 @@
|
|||||||
|
package nftables
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/nftables"
|
||||||
|
"github.com/google/nftables/expr"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TempDropOutputTCPRST temporarily drops outgoing TCP RST packets to the specified address and port,
|
||||||
|
// for any TCP packets not marked with the excludeMark given.
|
||||||
|
// This is necessary for TCP path MTU discovery to work, as the kernel will try to terminate the connection
|
||||||
|
// by sending a TCP RST packet, although we want to handle the connection manually.
|
||||||
|
func (f *Firewall) TempDropOutputTCPRST(_ context.Context,
|
||||||
|
src, dst netip.AddrPort, excludeMark int,
|
||||||
|
) (revert func(ctx context.Context) error, err error) {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
conn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating nftables connection: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
table, _, _, outputChain := setupFilterWithBaseChains(conn)
|
||||||
|
|
||||||
|
const maxExprsLen = 14
|
||||||
|
exprs := make([]expr.Any, 0, maxExprsLen)
|
||||||
|
|
||||||
|
// Match source IP
|
||||||
|
if src.Addr().Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 12, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: src.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 8, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: src.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match destination IP
|
||||||
|
if dst.Addr().Is4() {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 16, //nolint:mnd
|
||||||
|
Len: 4, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: dst.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseNetworkHeader,
|
||||||
|
Offset: 24, //nolint:mnd
|
||||||
|
Len: 16, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: dst.Addr().AsSlice(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match TCP protocol (6)
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||||
|
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, //nolint:mnd
|
||||||
|
)
|
||||||
|
|
||||||
|
// Match source port
|
||||||
|
srcPortBytes := []byte{byte(src.Port() >> 8), byte(src.Port())} //nolint:mnd
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 0, //nolint:mnd
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: srcPortBytes,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Match destination port
|
||||||
|
dstPortBytes := []byte{byte(dst.Port() >> 8), byte(dst.Port())} //nolint:mnd
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 2, //nolint:mnd
|
||||||
|
Len: 2, //nolint:mnd
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: dstPortBytes,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Match TCP RST flag (only RST set)
|
||||||
|
// TCP flags offset is 13th byte of the header (12 in 0-based)
|
||||||
|
// RST flag is bit 1 (value 0x04)
|
||||||
|
// We use bitwise to check mask == comparison (only RST is set)
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Payload{
|
||||||
|
DestRegister: 1,
|
||||||
|
Base: expr.PayloadBaseTransportHeader,
|
||||||
|
Offset: 13, //nolint:mnd
|
||||||
|
Len: 1,
|
||||||
|
},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpEq,
|
||||||
|
Register: 1,
|
||||||
|
Data: []byte{0x04}, //nolint:mnd
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Exclude packets with the mark using mark != excludeMark
|
||||||
|
markData := []byte{
|
||||||
|
byte(excludeMark), byte(excludeMark >> 8), byte(excludeMark >> 16), byte(excludeMark >> 24), //nolint:mnd
|
||||||
|
}
|
||||||
|
exprs = append(exprs,
|
||||||
|
&expr.Meta{Key: expr.MetaKeyMARK, Register: 1},
|
||||||
|
&expr.Cmp{
|
||||||
|
Op: expr.CmpOpNeq,
|
||||||
|
Register: 1,
|
||||||
|
Data: markData,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
// DROP verdict
|
||||||
|
exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop})
|
||||||
|
|
||||||
|
rule := &nftables.Rule{
|
||||||
|
Table: table,
|
||||||
|
Chain: outputChain,
|
||||||
|
Exprs: exprs,
|
||||||
|
}
|
||||||
|
|
||||||
|
conn.AddRule(rule)
|
||||||
|
f.rules = append(f.rules, rule)
|
||||||
|
|
||||||
|
err = conn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
f.rules = f.rules[:len(f.rules)-1]
|
||||||
|
return nil, fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capture rule for revert
|
||||||
|
ruleCopy := *rule
|
||||||
|
|
||||||
|
revert = func(_ context.Context) error {
|
||||||
|
f.mutex.Lock()
|
||||||
|
defer f.mutex.Unlock()
|
||||||
|
|
||||||
|
revertConn, err := nftables.New()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating nftables connection for revert: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = f.deleteRule(revertConn, &ruleCopy)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting rule: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = revertConn.Flush()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("flushing: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return revert, nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user