fix: add -w flag to iptables commands and env var for firewall backend

This commit is contained in:
Toby 2026-03-30 12:21:04 -07:00
parent 724e469f70
commit 25a2530ac7
2 changed files with 35 additions and 15 deletions

View file

@ -7,12 +7,16 @@ import (
"encoding/hex"
"fmt"
"net"
"os"
"os/exec"
"strings"
"sync"
eUtils "github.com/apernet/hysteria/extras/v2/utils"
)
const firewallBackendEnv = "HYSTERIA_FIREWALL_BACKEND"
type osCommandRunner struct{}
func (osCommandRunner) LookPath(file string) (string, error) {
@ -64,28 +68,39 @@ func setupUDPPortRedirectWithRunner(r commandRunner, listenAddr *net.UDPAddr, po
if len(redirects) == 0 {
return nil, nil
}
if _, err := r.LookPath("nft"); err == nil {
switch strings.ToLower(os.Getenv(firewallBackendEnv)) {
case "nftables", "nft":
return setupNFTablesRedirect(r, listenAddr, ports, redirects)
case "iptables", "ipt":
return setupIPTablesRedirect(r, listenAddr, ports, redirects)
default:
// Auto-detect: prefer nftables, fall back to iptables.
if _, err := r.LookPath("nft"); err == nil {
return setupNFTablesRedirect(r, listenAddr, ports, redirects)
}
return setupIPTablesRedirect(r, listenAddr, ports, redirects)
}
return setupIPTablesRedirect(r, listenAddr, ports, redirects)
}
func setupNFTablesRedirect(r commandRunner, listenAddr *net.UDPAddr, ports, redirects eUtils.PortUnion) (*closerFuncs, error) {
families := nftFamiliesForAddr(listenAddr)
cleanup := &closerFuncs{}
nft := func(args ...string) error {
return r.Run("nft", args...)
}
for _, family := range families {
hash := shortHash("nft|" + family + "|" + hashInput(listenAddr, ports))
tableName := "hysteria_" + hash
if err := r.Run("nft", "add", "table", family, tableName); err != nil {
if err := nft("add", "table", family, tableName); err != nil {
_ = cleanup.Close()
return nil, err
}
cleanup.add(func() { _ = r.Run("nft", "delete", "table", family, tableName) })
cleanup.add(func() { _ = nft("delete", "table", family, tableName) })
for _, chainArgs := range [][]string{
{"add", "chain", family, tableName, "prerouting", "{", "type", "nat", "hook", "prerouting", "priority", "dstnat;", "policy", "accept;", "}"},
{"add", "chain", family, tableName, "output", "{", "type", "nat", "hook", "output", "priority", "dstnat;", "policy", "accept;", "}"},
} {
if err := r.Run("nft", chainArgs...); err != nil {
if err := nft(chainArgs...); err != nil {
_ = cleanup.Close()
return nil, err
}
@ -97,7 +112,7 @@ func setupNFTablesRedirect(r commandRunner, listenAddr *net.UDPAddr, ports, redi
args = append(args, match...)
}
args = append(args, "udp", "dport", nftPortExpr(portRange), "redirect", "to", fmt.Sprintf(":%d", ports[0].Start))
if err := r.Run("nft", args...); err != nil {
if err := nft(args...); err != nil {
_ = cleanup.Close()
return nil, err
}
@ -114,18 +129,20 @@ func setupIPTablesRedirect(r commandRunner, listenAddr *net.UDPAddr, ports, redi
}
cleanup := &closerFuncs{}
for _, bin := range bins {
ipt := func(args ...string) error {
return r.Run(bin, append([]string{"-w"}, args...)...)
}
hash := shortHash(bin + "|" + hashInput(listenAddr, ports))
chainName := "HYSTERIA-PR-" + hash
if err := r.Run(bin, "-t", "nat", "-N", chainName); err != nil {
if err := ipt("-t", "nat", "-N", chainName); err != nil {
_ = cleanup.Close()
return nil, err
}
cleanup.add(func() {
_ = r.Run(bin, "-t", "nat", "-F", chainName)
_ = r.Run(bin, "-t", "nat", "-X", chainName)
_ = ipt("-t", "nat", "-F", chainName)
_ = ipt("-t", "nat", "-X", chainName)
})
redirectArgs := []string{"-t", "nat", "-A", chainName, "-p", "udp", "-j", "REDIRECT", "--to-ports", fmt.Sprintf("%d", ports[0].Start)}
if err := r.Run(bin, redirectArgs...); err != nil {
if err := ipt("-t", "nat", "-A", chainName, "-p", "udp", "-j", "REDIRECT", "--to-ports", fmt.Sprintf("%d", ports[0].Start)); err != nil {
_ = cleanup.Close()
return nil, err
}
@ -136,12 +153,12 @@ func setupIPTablesRedirect(r commandRunner, listenAddr *net.UDPAddr, ports, redi
args = append(args, match...)
}
args = append(args, "-p", "udp", "--dport", iptablesPortExpr(portRange), "-j", chainName)
if err := r.Run(bin, args...); err != nil {
if err := ipt(args...); err != nil {
_ = cleanup.Close()
return nil, err
}
deleteArgs := append([]string{"-t", "nat", "-D", baseChain}, args[4:]...)
cleanup.add(func() { _ = r.Run(bin, deleteArgs...) })
cleanup.add(func() { _ = ipt(deleteArgs...) })
}
}
}

View file

@ -68,8 +68,11 @@ func TestSetupUDPPortRedirectWithRunnerIPTablesFallback(t *testing.T) {
require.NoError(t, cleanup.Close())
foundDelete := false
for _, cmd := range runner.cmds {
if len(cmd) > 3 && cmd[2] == "-D" || (len(cmd) > 4 && cmd[3] == "-D") {
foundDelete = true
for _, arg := range cmd {
if arg == "-D" {
foundDelete = true
break
}
}
}
require.True(t, foundDelete)